Skip to main content

otf_pixels_codec_avif/av1/
transform.rs

1//! Inverse transforms and dequantization (spec §7.12–§7.13).
2//!
3//! A transform block is reconstructed by dequantizing the coefficient levels
4//! (`dequantize`), running a 2D inverse transform (`inverse_transform_2d`), and
5//! adding the result to the prediction (`add_residual`). The `dsp` submodule
6//! below holds the full lossy machinery — the DCT/ADST/identity butterfly
7//! network and the 2D driver — transcribed from the spec's ordered steps.
8//!
9//! The lossless path is just the qindex-0 corner of that machinery: `dc_q` and
10//! `ac_q` are both 4, `dqDenom` is 1, and the 4x4 Walsh–Hadamard transform
11//! divides the dequantiser's "times 4" straight back out, which is why lossless
12//! is bit-exact. It flows through the same `dequantize` + `inverse_transform_2d`
13//! pair as the lossy path, with `Lossless == 1`.
14
15/// Reconstruct a transform block by adding the `Residual` to the prediction and
16/// clipping to the sample range (`Clip1`, the reconstruct process §7.12.3 step
17/// 3), honouring the transform type's `flipUD`/`flipLR`. `prediction` is
18/// row-major `residual.width * residual.height`; the result is the same shape.
19///
20/// The residual is stored pre-flip, so `Residual[i][j]` is added to the output
21/// at `(flipUD ? h-1-i : i, flipLR ? w-1-j : j)` — the prediction already sits
22/// there. The lossless path is the `DCT_DCT` (no-flip) 4x4 corner of this.
23#[must_use]
24pub fn add_residual(
25    prediction: &[u16],
26    residual: &Residual,
27    tx_type: TxType,
28    bit_depth: u8,
29) -> Vec<u16> {
30    let w = residual.width;
31    let h = residual.height;
32    let max = i64::from((1_u32 << bit_depth) - 1);
33    let flip_ud = tx_type.flip_ud();
34    let flip_lr = tx_type.flip_lr();
35    let mut out = prediction.to_vec();
36    for i in 0..h {
37        for j in 0..w {
38            let xx = if flip_lr { w - j - 1 } else { j };
39            let yy = if flip_ud { h - i - 1 } else { i };
40            let idx = yy * w + xx;
41            let pred = prediction.get(idx).copied().unwrap_or(0);
42            let value = i64::from(pred) + i64::from(residual.at(i, j));
43            if let Some(cell) = out.get_mut(idx) {
44                *cell = value.clamp(0, max) as u16;
45            }
46        }
47    }
48    out
49}
50
51/// Reconstruct a 4x4 block (§7.12.3 step 3): a thin wrapper over [`add_residual`]
52/// for the `DCT_DCT` (no-flip) case the lossless 4x4 tile drives.
53#[must_use]
54pub fn add_residual_4x4(
55    prediction: &[[u16; 4]; 4],
56    residual: &Residual,
57    bit_depth: u8,
58) -> [[u16; 4]; 4] {
59    let mut flat = [0_u16; 16];
60    for (i, row) in prediction.iter().enumerate() {
61        for (j, &v) in row.iter().enumerate() {
62            if let Some(cell) = flat.get_mut(i * 4 + j) {
63                *cell = v;
64            }
65        }
66    }
67    let out = add_residual(&flat, residual, TxType::DctDct, bit_depth);
68    let mut result = [[0_u16; 4]; 4];
69    for (i, row) in result.iter_mut().enumerate() {
70        for (j, cell) in row.iter_mut().enumerate() {
71            *cell = out.get(i * 4 + j).copied().unwrap_or(0);
72        }
73    }
74    result
75}
76
77pub use dsp::{
78    Dequant, Residual, TxSize, TxType, ac_q, dc_q, dequantize, dequantize_with_matrix,
79    inverse_transform_2d, quantizer_matrix,
80};
81
82/// The lossy inverse-transform machinery (spec §7.13): the DCT/ADST/identity
83/// butterfly network, the 2D transform driver, and the dequantiser lookups.
84///
85/// This lives in its own module because the 1D transforms are transcribed
86/// straight from the spec's ordered `B`/`H` butterfly steps, which read and
87/// write a fixed-length working array `T` by index. Expressing that through
88/// `.get()` would obscure the correspondence with the spec, so the module opts
89/// into `indexing_slicing`: every index is a spec constant bounded below the
90/// array length (`T` is 64 long, the largest transform).
91pub(crate) mod dsp {
92    #![allow(
93        clippy::indexing_slicing,
94        clippy::needless_range_loop,
95        reason = "fixed-size DSP working arrays; every index is a spec-bounded \
96                  constant strictly below the array length, and the permutation \
97                  loops index a separate snapshot from the array they write"
98    )]
99
100    /// `Round2(x, n)` (§4.7): rounding right shift, `x` unchanged when `n == 0`.
101    fn round2(x: i64, n: u32) -> i64 {
102        if n == 0 { x } else { (x + (1 << (n - 1))) >> n }
103    }
104
105    /// `Clip3(-(1<<(bits-1)), (1<<(bits-1))-1, x)`: clamp to a signed `bits` range.
106    fn clip_signed(x: i64, bits: u32) -> i64 {
107        let lo = -(1_i64 << (bits - 1));
108        let hi = (1_i64 << (bits - 1)) - 1;
109        x.clamp(lo, hi)
110    }
111
112    /// `Cos128_Lookup` (§7.13.2.1): `round(4096 * cos(angle * pi / 128))` for
113    /// `angle` in `0..=64`.
114    const COS128_LOOKUP: [i64; 65] = [
115        4096, 4095, 4091, 4085, 4076, 4065, 4052, 4036, 4017, 3996, 3973, 3948, 3920, 3889, 3857,
116        3822, 3784, 3745, 3703, 3659, 3612, 3564, 3513, 3461, 3406, 3349, 3290, 3229, 3166, 3102,
117        3035, 2967, 2896, 2824, 2751, 2675, 2598, 2520, 2440, 2359, 2276, 2191, 2106, 2019, 1931,
118        1842, 1751, 1660, 1567, 1474, 1380, 1285, 1189, 1092, 995, 897, 799, 700, 601, 501, 401,
119        301, 201, 101, 0,
120    ];
121
122    /// `cos128(angle)` (§7.13.2.1), reducing the angle modulo 256.
123    fn cos128(angle: i32) -> i64 {
124        let a = (angle & 255) as usize;
125        match a {
126            0..=64 => COS128_LOOKUP[a],
127            65..=128 => -COS128_LOOKUP[128 - a],
128            129..=192 => -COS128_LOOKUP[a - 128],
129            _ => COS128_LOOKUP[256 - a],
130        }
131    }
132
133    /// `sin128(angle) = cos128(angle - 64)` (§7.13.2.1).
134    fn sin128(angle: i32) -> i64 {
135        cos128(angle - 64)
136    }
137
138    /// `brev(numBits, x)` (§7.13.2.1): bit-reversal of the low `num_bits` of `x`.
139    fn brev(num_bits: u32, x: usize) -> usize {
140        let mut t = 0;
141        for i in 0..num_bits {
142            let bit = (x >> i) & 1;
143            t += bit << (num_bits - 1 - i);
144        }
145        t
146    }
147
148    /// `B(a, b, angle, flip, r)`: butterfly rotation (§7.13.2.1). The `r` clamp is
149    /// a conformance requirement on the inputs, not an operation, so it is unused.
150    fn butterfly(t: &mut [i64; 64], a: usize, b: usize, angle: i32, flip: bool) {
151        let (ta, tb) = (t[a], t[b]);
152        let x = ta * cos128(angle) - tb * sin128(angle);
153        let y = ta * sin128(angle) + tb * cos128(angle);
154        t[a] = round2(x, 12);
155        t[b] = round2(y, 12);
156        if flip {
157            t.swap(a, b);
158        }
159    }
160
161    /// `H(a, b, flip, r)`: Hadamard rotation (§7.13.2.1). A flip swaps the pair.
162    fn hadamard(t: &mut [i64; 64], a: usize, b: usize, flip: bool, r: u32) {
163        let (a, b) = if flip { (b, a) } else { (a, b) };
164        let (x, y) = (t[a], t[b]);
165        t[a] = clip_signed(x + y, r);
166        t[b] = clip_signed(x - y, r);
167    }
168
169    /// Inverse DCT array permutation process (§7.13.2.2).
170    fn dct_permute(t: &mut [i64; 64], n: u32) {
171        let len = 1usize << n;
172        let copy = *t;
173        for i in 0..len {
174            t[i] = copy[brev(n, i)];
175        }
176    }
177
178    /// Inverse DCT process (§7.13.2.3) for a length-`2^n` array (`2 <= n <= 6`).
179    #[allow(
180        clippy::too_many_lines,
181        reason = "a faithful transcription of the spec's 31 ordered butterfly steps"
182    )]
183    fn inverse_dct(t: &mut [i64; 64], n: u32, r: u32) {
184        dct_permute(t, n);
185        // Each numbered block matches the like-numbered step of §7.13.2.3.
186        if n == 6 {
187            for i in 0..16 {
188                butterfly(t, 32 + i, 63 - i, 63 - 4 * brev(4, i) as i32, false);
189            }
190        }
191        if n >= 5 {
192            for i in 0..8 {
193                butterfly(t, 16 + i, 31 - i, 6 + ((brev(3, 7 - i) as i32) << 3), false);
194            }
195        }
196        if n == 6 {
197            for i in 0..16 {
198                hadamard(t, 32 + i * 2, 33 + i * 2, i & 1 == 1, r);
199            }
200        }
201        if n >= 4 {
202            for i in 0..4 {
203                butterfly(t, 8 + i, 15 - i, 12 + ((brev(2, 3 - i) as i32) << 4), false);
204            }
205        }
206        if n >= 5 {
207            for i in 0..8 {
208                hadamard(t, 16 + 2 * i, 17 + 2 * i, i & 1 == 1, r);
209            }
210        }
211        if n == 6 {
212            for i in 0..4 {
213                for j in 0..2 {
214                    butterfly(
215                        t,
216                        62 - i * 4 - j,
217                        33 + i * 4 + j,
218                        60 - 16 * brev(2, i) as i32 + 64 * j as i32,
219                        true,
220                    );
221                }
222            }
223        }
224        if n >= 3 {
225            for i in 0..2 {
226                butterfly(t, 4 + i, 7 - i, 56 - 32 * i as i32, false);
227            }
228        }
229        if n >= 4 {
230            for i in 0..4 {
231                hadamard(t, 8 + 2 * i, 9 + 2 * i, i & 1 == 1, r);
232            }
233        }
234        if n >= 5 {
235            for i in 0..2 {
236                for j in 0..2 {
237                    butterfly(
238                        t,
239                        30 - 4 * i - j,
240                        17 + 4 * i + j,
241                        24 + ((j as i32) << 6) + (((1 - i) as i32) << 5),
242                        true,
243                    );
244                }
245            }
246        }
247        if n == 6 {
248            for i in 0..8 {
249                for j in 0..2 {
250                    hadamard(t, 32 + i * 4 + j, 35 + i * 4 - j, i & 1 == 1, r);
251                }
252            }
253        }
254        for i in 0..2 {
255            butterfly(t, 2 * i, 2 * i + 1, 32 + 16 * i as i32, i == 0);
256        }
257        if n >= 3 {
258            for i in 0..2 {
259                hadamard(t, 4 + 2 * i, 5 + 2 * i, i == 1, r);
260            }
261        }
262        if n >= 4 {
263            for i in 0..2 {
264                butterfly(t, 14 - i, 9 + i, 48 + 64 * i as i32, true);
265            }
266        }
267        if n >= 5 {
268            for i in 0..4 {
269                for j in 0..2 {
270                    hadamard(t, 16 + 4 * i + j, 19 + 4 * i - j, i & 1 == 1, r);
271                }
272            }
273        }
274        if n == 6 {
275            for i in 0..2 {
276                for j in 0..4 {
277                    butterfly(
278                        t,
279                        61 - i * 8 - j,
280                        34 + i * 8 + j,
281                        56 - i as i32 * 32 + (j as i32 >> 1) * 64,
282                        true,
283                    );
284                }
285            }
286        }
287        for i in 0..2 {
288            hadamard(t, i, 3 - i, false, r);
289        }
290        if n >= 3 {
291            butterfly(t, 6, 5, 32, true);
292        }
293        if n >= 4 {
294            for i in 0..2 {
295                for j in 0..2 {
296                    hadamard(t, 8 + 4 * i + j, 11 + 4 * i - j, i == 1, r);
297                }
298            }
299        }
300        if n >= 5 {
301            for i in 0..4 {
302                butterfly(t, 29 - i, 18 + i, 48 + (i as i32 >> 1) * 64, true);
303            }
304        }
305        if n == 6 {
306            for i in 0..4 {
307                for j in 0..4 {
308                    hadamard(t, 32 + 8 * i + j, 39 + 8 * i - j, i & 1 == 1, r);
309                }
310            }
311        }
312        if n >= 3 {
313            for i in 0..4 {
314                hadamard(t, i, 7 - i, false, r);
315            }
316        }
317        if n >= 4 {
318            for i in 0..2 {
319                butterfly(t, 13 - i, 10 + i, 32, true);
320            }
321        }
322        if n >= 5 {
323            for i in 0..2 {
324                for j in 0..4 {
325                    hadamard(t, 16 + i * 8 + j, 23 + i * 8 - j, i == 1, r);
326                }
327            }
328        }
329        if n == 6 {
330            for i in 0..8 {
331                butterfly(t, 59 - i, 36 + i, if i < 4 { 48 } else { 112 }, true);
332            }
333        }
334        if n >= 4 {
335            for i in 0..8 {
336                hadamard(t, i, 15 - i, false, r);
337            }
338        }
339        if n >= 5 {
340            for i in 0..4 {
341                butterfly(t, 27 - i, 20 + i, 32, true);
342            }
343        }
344        if n == 6 {
345            for i in 0..8 {
346                hadamard(t, 32 + i, 47 - i, false, r);
347                hadamard(t, 48 + i, 63 - i, true, r);
348            }
349        }
350        if n >= 5 {
351            for i in 0..16 {
352                hadamard(t, i, 31 - i, false, r);
353            }
354        }
355        if n == 6 {
356            for i in 0..8 {
357                butterfly(t, 55 - i, 40 + i, 32, true);
358            }
359        }
360        if n == 6 {
361            for i in 0..32 {
362                hadamard(t, i, 63 - i, false, r);
363            }
364        }
365    }
366
367    /// Inverse ADST input array permutation (§7.13.2.4), `3 <= n <= 4`.
368    fn adst_permute_in(t: &mut [i64; 64], n: u32) {
369        let n0 = 1usize << n;
370        let copy = *t;
371        for i in 0..n0 {
372            let idx = if i & 1 == 1 { i - 1 } else { n0 - i - 1 };
373            t[i] = copy[idx];
374        }
375    }
376
377    /// Inverse ADST output array permutation (§7.13.2.5), `3 <= n <= 4`.
378    fn adst_permute_out(t: &mut [i64; 64], n: u32) {
379        let n0 = 1usize << n;
380        let copy = *t;
381        for i in 0..n0 {
382            let a = (i >> 3) & 1;
383            let b = ((i >> 2) & 1) ^ ((i >> 3) & 1);
384            let c = ((i >> 1) & 1) ^ ((i >> 2) & 1);
385            let d = (i & 1) ^ ((i >> 1) & 1);
386            let idx = ((d << 3) | (c << 2) | (b << 1) | a) >> (4 - n);
387            t[i] = if i & 1 == 1 { -copy[idx] } else { copy[idx] };
388        }
389    }
390
391    /// Inverse ADST4 process (§7.13.2.6).
392    fn inverse_adst4(t: &mut [i64; 64]) {
393        const SINPI_1_9: i64 = 1321;
394        const SINPI_2_9: i64 = 2482;
395        const SINPI_3_9: i64 = 3344;
396        const SINPI_4_9: i64 = 3803;
397        let (t0, t1, t2, t3) = (t[0], t[1], t[2], t[3]);
398        let mut s = [
399            SINPI_1_9 * t0,
400            SINPI_2_9 * t0,
401            SINPI_3_9 * t1,
402            SINPI_4_9 * t2,
403            SINPI_1_9 * t2,
404            SINPI_2_9 * t3,
405            SINPI_4_9 * t3,
406        ];
407        let a7 = t0 - t2;
408        let b7 = a7 + t3;
409        s[0] += s[3];
410        s[1] -= s[4];
411        s[3] = s[2];
412        s[2] = SINPI_3_9 * b7;
413        s[0] += s[5];
414        s[1] -= s[6];
415        let x0 = s[0] + s[3];
416        let x1 = s[1] + s[3];
417        let x2 = s[2];
418        let x3 = s[0] + s[1] - s[3];
419        t[0] = round2(x0, 12);
420        t[1] = round2(x1, 12);
421        t[2] = round2(x2, 12);
422        t[3] = round2(x3, 12);
423    }
424
425    /// Inverse ADST8 process (§7.13.2.7).
426    fn inverse_adst8(t: &mut [i64; 64], r: u32) {
427        adst_permute_in(t, 3);
428        for i in 0..4 {
429            butterfly(t, 2 * i, 2 * i + 1, 60 - 16 * i as i32, true);
430        }
431        for i in 0..4 {
432            hadamard(t, i, 4 + i, false, r);
433        }
434        for i in 0..2 {
435            butterfly(t, 4 + 3 * i, 5 + i, 48 - 32 * i as i32, true);
436        }
437        for i in 0..2 {
438            for j in 0..2 {
439                hadamard(t, 4 * j + i, 2 + 4 * j + i, false, r);
440            }
441        }
442        for i in 0..2 {
443            butterfly(t, 2 + 4 * i, 3 + 4 * i, 32, true);
444        }
445        adst_permute_out(t, 3);
446    }
447
448    /// Inverse ADST16 process (§7.13.2.8).
449    fn inverse_adst16(t: &mut [i64; 64], r: u32) {
450        adst_permute_in(t, 4);
451        for i in 0..8 {
452            butterfly(t, 2 * i, 2 * i + 1, 62 - 8 * i as i32, true);
453        }
454        for i in 0..8 {
455            hadamard(t, i, 8 + i, false, r);
456        }
457        for i in 0..2 {
458            butterfly(t, 8 + 2 * i, 9 + 2 * i, 56 - 32 * i as i32, true);
459            butterfly(t, 13 + 2 * i, 12 + 2 * i, 8 + 32 * i as i32, true);
460        }
461        for i in 0..4 {
462            for j in 0..2 {
463                hadamard(t, 8 * j + i, 4 + 8 * j + i, false, r);
464            }
465        }
466        for i in 0..2 {
467            for j in 0..2 {
468                butterfly(
469                    t,
470                    4 + 8 * j + 3 * i,
471                    5 + 8 * j + i,
472                    48 - 32 * i as i32,
473                    true,
474                );
475            }
476        }
477        for i in 0..2 {
478            for j in 0..4 {
479                hadamard(t, 4 * j + i, 2 + 4 * j + i, false, r);
480            }
481        }
482        for i in 0..4 {
483            butterfly(t, 2 + 4 * i, 3 + 4 * i, 32, true);
484        }
485        adst_permute_out(t, 4);
486    }
487
488    /// Inverse ADST process (§7.13.2.9) for a length-`2^n` array (`2 <= n <= 4`).
489    fn inverse_adst(t: &mut [i64; 64], n: u32, r: u32) {
490        match n {
491            2 => inverse_adst4(t),
492            3 => inverse_adst8(t, r),
493            _ => inverse_adst16(t, r),
494        }
495    }
496
497    /// Inverse identity transform process (§7.13.2.15), `2 <= n <= 5`.
498    fn inverse_identity(t: &mut [i64; 64], n: u32) {
499        let len = 1usize << n;
500        for cell in t.iter_mut().take(len) {
501            *cell = match n {
502                2 => round2(*cell * 5793, 12),
503                3 => *cell * 2,
504                4 => round2(*cell * 11586, 12),
505                _ => *cell * 4,
506            };
507        }
508    }
509
510    /// Inverse Walsh–Hadamard transform (§7.13.2.10) on the first four elements.
511    fn inverse_wht(t: &mut [i64; 64], shift: u32) {
512        let mut a = t[0] >> shift;
513        let mut c = t[1] >> shift;
514        let mut d = t[2] >> shift;
515        let mut b = t[3] >> shift;
516        a += c;
517        d -= b;
518        let e = (a - d) >> 1;
519        b = e - b;
520        c = e - c;
521        a -= b;
522        d += c;
523        t[0] = a;
524        t[1] = b;
525        t[2] = c;
526        t[3] = d;
527    }
528
529    /// Which 1D transform a row or column pass runs, before flipping.
530    #[derive(Clone, Copy, PartialEq, Eq)]
531    enum Kind {
532        Dct,
533        Adst,
534        Identity,
535    }
536
537    fn apply_1d(t: &mut [i64; 64], kind: Kind, n: u32, r: u32) {
538        match kind {
539            Kind::Dct => inverse_dct(t, n, r),
540            Kind::Adst => inverse_adst(t, n, r),
541            Kind::Identity => inverse_identity(t, n),
542        }
543    }
544
545    /// The 19 transform sizes (`TX_SIZES_ALL`, §6.10.28 order).
546    #[derive(Clone, Copy, PartialEq, Eq, Debug)]
547    #[allow(missing_docs, reason = "each variant is a self-describing WxH size")]
548    pub enum TxSize {
549        Tx4x4,
550        Tx8x8,
551        Tx16x16,
552        Tx32x32,
553        Tx64x64,
554        Tx4x8,
555        Tx8x4,
556        Tx8x16,
557        Tx16x8,
558        Tx16x32,
559        Tx32x16,
560        Tx32x64,
561        Tx64x32,
562        Tx4x16,
563        Tx16x4,
564        Tx8x32,
565        Tx32x8,
566        Tx16x64,
567        Tx64x16,
568    }
569
570    const TX_WIDTH_LOG2: [u32; 19] = [2, 3, 4, 5, 6, 2, 3, 3, 4, 4, 5, 5, 6, 2, 4, 3, 5, 4, 6];
571    const TX_HEIGHT_LOG2: [u32; 19] = [2, 3, 4, 5, 6, 3, 2, 4, 3, 5, 4, 6, 5, 4, 2, 5, 3, 6, 4];
572    const TRANSFORM_ROW_SHIFT: [u32; 19] =
573        [0, 1, 2, 2, 2, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2];
574    // `Tx_Size_Sqr[txSz]` / `Tx_Size_Sqr_Up[txSz]` as the square size's index
575    // (TX_4X4=0..TX_64X64=4): the square tx with side Min(w,h) resp. Max(w,h).
576    const TX_SIZE_SQR: [u32; 19] = [0, 1, 2, 3, 4, 0, 0, 1, 1, 2, 2, 3, 3, 0, 0, 1, 1, 2, 2];
577    const TX_SIZE_SQR_UP: [u32; 19] = [0, 1, 2, 3, 4, 1, 1, 2, 2, 3, 3, 4, 4, 2, 2, 3, 3, 4, 4];
578    // `Adjusted_Tx_Size[txSz]`: the coded tx size, mapping any 64-wide/high size
579    // down to its 32 counterpart (only the top-left 32x32 of coeffs is coded).
580    const ADJUSTED_TX_SIZE: [usize; 19] = [
581        0, 1, 2, 3, 3, 5, 6, 7, 8, 9, 10, 3, 3, 13, 14, 15, 16, 9, 10,
582    ];
583
584    /// All 19 transform sizes in `TX_SIZES_ALL` order, for index conversion.
585    const ALL_TX_SIZES: [TxSize; 19] = [
586        TxSize::Tx4x4,
587        TxSize::Tx8x8,
588        TxSize::Tx16x16,
589        TxSize::Tx32x32,
590        TxSize::Tx64x64,
591        TxSize::Tx4x8,
592        TxSize::Tx8x4,
593        TxSize::Tx8x16,
594        TxSize::Tx16x8,
595        TxSize::Tx16x32,
596        TxSize::Tx32x16,
597        TxSize::Tx32x64,
598        TxSize::Tx64x32,
599        TxSize::Tx4x16,
600        TxSize::Tx16x4,
601        TxSize::Tx8x32,
602        TxSize::Tx32x8,
603        TxSize::Tx16x64,
604        TxSize::Tx64x16,
605    ];
606
607    impl TxSize {
608        /// The transform size for a `TX_SIZES_ALL` index, or `TX_4X4` if out of
609        /// range. Inverse of `self as usize`.
610        #[must_use]
611        pub fn from_index(index: usize) -> TxSize {
612            ALL_TX_SIZES.get(index).copied().unwrap_or(TxSize::Tx4x4)
613        }
614
615        /// `Tx_Width_Log2[txSz]`.
616        #[must_use]
617        pub fn log2_width(self) -> u32 {
618            TX_WIDTH_LOG2[self as usize]
619        }
620
621        /// `Tx_Height_Log2[txSz]`.
622        #[must_use]
623        pub fn log2_height(self) -> u32 {
624            TX_HEIGHT_LOG2[self as usize]
625        }
626
627        /// `Tx_Size_Sqr[txSz]` as the square size's index (`Min(w,h)`).
628        #[must_use]
629        pub fn sqr_idx(self) -> u32 {
630            TX_SIZE_SQR[self as usize]
631        }
632
633        /// `Tx_Size_Sqr_Up[txSz]` as the square size's index (`Max(w,h)`).
634        #[must_use]
635        pub fn sqr_up_idx(self) -> u32 {
636            TX_SIZE_SQR_UP[self as usize]
637        }
638
639        /// `txSzCtx = (Tx_Size_Sqr + Tx_Size_Sqr_Up + 1) >> 1`, the coefficient
640        /// CDF bucket (0..=4).
641        #[must_use]
642        pub fn tx_size_ctx(self) -> usize {
643            ((self.sqr_idx() + self.sqr_up_idx() + 1) >> 1) as usize
644        }
645
646        /// `Tx_Width_Log2[Adjusted_Tx_Size[txSz]]`: the coded block's width log2,
647        /// which drives the coefficient position maths (`bwl`).
648        #[must_use]
649        pub fn adjusted_log2_width(self) -> u32 {
650            TX_WIDTH_LOG2[ADJUSTED_TX_SIZE[self as usize]]
651        }
652
653        /// `Tx_Height[Adjusted_Tx_Size[txSz]]`: the coded block's height.
654        #[must_use]
655        pub fn adjusted_height(self) -> usize {
656            1 << TX_HEIGHT_LOG2[ADJUSTED_TX_SIZE[self as usize]]
657        }
658
659        /// `Tx_Width[Adjusted_Tx_Size[txSz]]`: the coded block's width.
660        #[must_use]
661        pub fn adjusted_width(self) -> usize {
662            1 << self.adjusted_log2_width()
663        }
664
665        /// `segEob` (§5.11.39): the number of scan positions the block can code.
666        #[must_use]
667        pub fn seg_eob(self) -> usize {
668            match self {
669                TxSize::Tx16x64 | TxSize::Tx64x16 => 512,
670                _ => (self.width() * self.height()).min(1024),
671            }
672        }
673
674        /// `eobMultisize` (§5.11.39): selects the `eob_pt_*` alphabet (0..=6).
675        #[must_use]
676        pub fn eob_multisize(self) -> usize {
677            (self.log2_width().min(5) + self.log2_height().min(5) - 4) as usize
678        }
679
680        /// The transform width in samples, `1 << Tx_Width_Log2[txSz]`.
681        #[must_use]
682        pub fn width(self) -> usize {
683            1 << self.log2_width()
684        }
685
686        /// The transform height in samples, `1 << Tx_Height_Log2[txSz]`.
687        #[must_use]
688        pub fn height(self) -> usize {
689            1 << self.log2_height()
690        }
691
692        fn row_shift(self) -> u32 {
693            TRANSFORM_ROW_SHIFT[self as usize]
694        }
695
696        /// `dqDenom` (§7.12.3): the shared denominator of the dequantiser.
697        fn dq_denom(self) -> i64 {
698            match self {
699                TxSize::Tx32x32
700                | TxSize::Tx16x32
701                | TxSize::Tx32x16
702                | TxSize::Tx16x64
703                | TxSize::Tx64x16 => 2,
704                TxSize::Tx64x64 | TxSize::Tx32x64 | TxSize::Tx64x32 => 4,
705                _ => 1,
706            }
707        }
708    }
709
710    /// `PlaneTxType`: the 16 transform types (§6.10.28 order). The first word
711    /// names the column (vertical) transform, the second the row (horizontal).
712    #[derive(Clone, Copy, PartialEq, Eq, Debug)]
713    #[allow(
714        missing_docs,
715        reason = "each variant names its column_row transform pair per §6.10.28"
716    )]
717    pub enum TxType {
718        DctDct,
719        AdstDct,
720        DctAdst,
721        AdstAdst,
722        FlipadstDct,
723        DctFlipadst,
724        FlipadstFlipadst,
725        AdstFlipadst,
726        FlipadstAdst,
727        Idtx,
728        VDct,
729        HDct,
730        VAdst,
731        HAdst,
732        VFlipadst,
733        HFlipadst,
734    }
735
736    impl TxType {
737        /// The row (horizontal) 1D transform — the type's second word.
738        fn row_kind(self) -> Kind {
739            match self {
740                TxType::DctDct | TxType::AdstDct | TxType::FlipadstDct | TxType::HDct => Kind::Dct,
741                TxType::Idtx | TxType::VDct | TxType::VAdst | TxType::VFlipadst => Kind::Identity,
742                _ => Kind::Adst,
743            }
744        }
745
746        /// The column (vertical) 1D transform — the type's first word.
747        fn col_kind(self) -> Kind {
748            match self {
749                TxType::DctDct | TxType::DctAdst | TxType::DctFlipadst | TxType::VDct => Kind::Dct,
750                TxType::Idtx | TxType::HDct | TxType::HAdst | TxType::HFlipadst => Kind::Identity,
751                _ => Kind::Adst,
752            }
753        }
754
755        /// `flipUD` (§7.12.3): the column transform is a flipped ADST.
756        #[must_use]
757        pub fn flip_ud(self) -> bool {
758            matches!(
759                self,
760                TxType::FlipadstDct
761                    | TxType::FlipadstAdst
762                    | TxType::VFlipadst
763                    | TxType::FlipadstFlipadst
764            )
765        }
766
767        /// `flipLR` (§7.12.3): the row transform is a flipped ADST.
768        #[must_use]
769        pub fn flip_lr(self) -> bool {
770            matches!(
771                self,
772                TxType::DctFlipadst
773                    | TxType::AdstFlipadst
774                    | TxType::HFlipadst
775                    | TxType::FlipadstFlipadst
776            )
777        }
778    }
779
780    /// A reconstructed residual block, `width` by `height` samples in raster
781    /// order. Values are pre-flip: the caller applies `flip_ud`/`flip_lr` when
782    /// adding to the prediction (§7.12.3 step 3).
783    pub struct Residual {
784        /// The residual width in samples.
785        pub width: usize,
786        /// The residual height in samples.
787        pub height: usize,
788        values: [i32; 64 * 64],
789    }
790
791    impl Residual {
792        /// The residual at row `i`, column `j` (0 outside the block).
793        #[must_use]
794        pub fn at(&self, i: usize, j: usize) -> i32 {
795            if i < self.height && j < self.width {
796                self.values[i * self.width + j]
797            } else {
798                0
799            }
800        }
801    }
802
803    /// 2D inverse transform process (§7.13.3). `dequant` holds the dequantised
804    /// coefficients `Dequant[i][j]` in raster order over the `min(32,w)` by
805    /// `min(32,h)` populated region; entries beyond it are treated as zero.
806    #[must_use]
807    pub fn inverse_transform_2d(
808        dequant: &Dequant,
809        tx_size: TxSize,
810        tx_type: TxType,
811        lossless: bool,
812        bit_depth: u8,
813    ) -> Residual {
814        let log2w = tx_size.log2_width();
815        let log2h = tx_size.log2_height();
816        let w = 1usize << log2w;
817        let h = 1usize << log2h;
818        let row_shift = if lossless { 0 } else { tx_size.row_shift() };
819        let col_shift = if lossless { 0 } else { 4 };
820        let row_clamp = u32::from(bit_depth) + 8;
821        let col_clamp = (u32::from(bit_depth) + 6).max(16);
822        let rect_scale = log2w.abs_diff(log2h) == 1;
823
824        let mut residual = [0_i64; 64 * 64];
825        let mut t = [0_i64; 64];
826
827        // Row transforms.
828        for i in 0..h {
829            for (j, cell) in t.iter_mut().enumerate().take(w) {
830                *cell = if i < 32 && j < 32 {
831                    dequant.at(i, j)
832                } else {
833                    0
834                };
835            }
836            if rect_scale {
837                for cell in t.iter_mut().take(w) {
838                    *cell = round2(*cell * 2896, 12);
839                }
840            }
841            if lossless {
842                inverse_wht(&mut t, 2);
843            } else {
844                apply_1d(&mut t, tx_type.row_kind(), log2w, row_clamp);
845            }
846            for j in 0..w {
847                residual[i * w + j] = round2(t[j], row_shift);
848            }
849        }
850
851        // Clamp between the passes.
852        for value in residual.iter_mut().take(w * h) {
853            *value = clip_signed(*value, col_clamp);
854        }
855
856        // Column transforms.
857        for j in 0..w {
858            for (i, cell) in t.iter_mut().enumerate().take(h) {
859                *cell = residual[i * w + j];
860            }
861            if lossless {
862                inverse_wht(&mut t, 0);
863            } else {
864                apply_1d(&mut t, tx_type.col_kind(), log2h, col_clamp);
865            }
866            for i in 0..h {
867                residual[i * w + j] = round2(t[i], col_shift);
868            }
869        }
870
871        let mut values = [0_i32; 64 * 64];
872        for (out, &v) in values.iter_mut().zip(residual.iter()).take(w * h) {
873            *out = v as i32;
874        }
875        Residual {
876            width: w,
877            height: h,
878            values,
879        }
880    }
881
882    /// The forward transform the encoder needs, derived from the inverse one
883    /// above rather than transcribed: AV1's inverse transform is separable
884    /// (rows, a shift, columns, a shift), so probing each 1D inverse with
885    /// impulses gives its synthesis matrix, and the forward transform is the
886    /// projection onto those near-orthogonal bases. Scaling then matches the
887    /// decoder by construction, and only the inverse is normative.
888    #[derive(Debug, Clone)]
889    pub(crate) struct ForwardBasis {
890        w: usize,
891        h: usize,
892        /// `row[x * w + j]`: output sample `x` of the row inverse of impulse `j`.
893        row: Vec<f64>,
894        /// `col[y * h + i]`: the column inverse, likewise.
895        col: Vec<f64>,
896        /// Squared norm of each basis vector.
897        row_norm: Vec<f64>,
898        col_norm: Vec<f64>,
899        /// The shifts and rectangular scaling between the two passes.
900        gain: f64,
901        flip_ud: bool,
902        flip_lr: bool,
903    }
904
905    fn basis_1d(kind: Kind, log2n: u32) -> (Vec<f64>, Vec<f64>) {
906        const IMPULSE: i64 = 1 << 12;
907        let n = 1usize << log2n;
908        let mut m = vec![0.0; n * n];
909        for j in 0..n {
910            let mut t = [0_i64; 64];
911            t[j] = IMPULSE;
912            apply_1d(&mut t, kind, log2n, 40);
913            for x in 0..n {
914                m[x * n + j] = t[x] as f64 / IMPULSE as f64;
915            }
916        }
917        let norm = (0..n)
918            .map(|j| (0..n).map(|x| m[x * n + j] * m[x * n + j]).sum())
919            .collect();
920        (m, norm)
921    }
922
923    impl ForwardBasis {
924        /// The basis for a lossy (non-WHT) transform of `tx_size` and `tx_type`.
925        pub(crate) fn new(tx_size: TxSize, tx_type: TxType) -> Self {
926            let (log2w, log2h) = (tx_size.log2_width(), tx_size.log2_height());
927            let (row, row_norm) = basis_1d(tx_type.row_kind(), log2w);
928            let (col, col_norm) = basis_1d(tx_type.col_kind(), log2h);
929            let rect = if log2w.abs_diff(log2h) == 1 {
930                2896.0 / 4096.0
931            } else {
932                1.0
933            };
934            let gain = rect / f64::from(1_u32 << (tx_size.row_shift() + 4));
935            Self {
936                w: 1 << log2w,
937                h: 1 << log2h,
938                row,
939                col,
940                row_norm,
941                col_norm,
942                gain,
943                flip_ud: tx_type.flip_ud(),
944                flip_lr: tx_type.flip_lr(),
945            }
946        }
947
948        /// Coefficients (in the dequantised domain the inverse takes, raster
949        /// order) whose inverse transform is `residual` (`w * h`, row-major).
950        pub(crate) fn forward(&self, residual: &[i32]) -> Vec<f64> {
951            let (w, h) = (self.w, self.h);
952            let at = |y: usize, x: usize| {
953                let y = if self.flip_ud { h - 1 - y } else { y };
954                let x = if self.flip_lr { w - 1 - x } else { x };
955                f64::from(residual[y * w + x])
956            };
957            // Rows: project each residual row onto the row basis.
958            let mut tmp = vec![0.0; w * h];
959            for y in 0..h {
960                for j in 0..w {
961                    let mut acc = 0.0;
962                    for x in 0..w {
963                        acc += at(y, x) * self.row[x * w + j];
964                    }
965                    tmp[y * w + j] = acc / self.row_norm[j];
966                }
967            }
968            // Columns.
969            let mut out = vec![0.0; w * h];
970            for i in 0..h {
971                for j in 0..w {
972                    let mut acc = 0.0;
973                    for y in 0..h {
974                        acc += tmp[y * w + j] * self.col[y * h + i];
975                    }
976                    out[i * w + j] = acc / self.col_norm[i] / self.gain;
977                }
978            }
979            out
980        }
981    }
982
983    /// The level the dequantiser turns into (about) `coefficient`, with a
984    /// dead zone: `bias` is the rounding point as a fraction of a step.
985    pub(crate) fn quantize(coefficient: f64, q: i64, tx_size: TxSize, bias: f64) -> i32 {
986        let scaled = coefficient.abs() * tx_size.dq_denom() as f64 / q as f64;
987        let level = (scaled + bias).floor().min(f64::from(1 << 20)) as i32;
988        if coefficient < 0.0 { -level } else { level }
989    }
990
991    /// The dequantised coefficient block `Dequant[i][j]`, raster order over the
992    /// populated `tw` by `th` region (`tw = min(32,w)`, `th = min(32,h)`).
993    pub struct Dequant {
994        width: usize,
995        height: usize,
996        values: [i64; 32 * 32],
997    }
998
999    impl Dequant {
1000        fn at(&self, i: usize, j: usize) -> i64 {
1001            if i < self.height && j < self.width {
1002                self.values[i * self.width + j]
1003            } else {
1004                0
1005            }
1006        }
1007    }
1008
1009    /// Dequantise one transform block (§7.12.3 step 1). `quant` holds `Quant[]`
1010    /// in raster order over the `tw` by `th` region; `dc_quant`/`ac_quant` are
1011    /// the plane's DC/AC quantiser steps. No quantiser matrix is applied.
1012    #[must_use]
1013    pub fn dequantize(
1014        quant: &[i32],
1015        tx_size: TxSize,
1016        dc_quant: i64,
1017        ac_quant: i64,
1018        bit_depth: u8,
1019    ) -> Dequant {
1020        dequantize_with_matrix(quant, tx_size, dc_quant, ac_quant, None, bit_depth)
1021    }
1022
1023    /// [`dequantize`], with each position's quantizer first weighted by a
1024    /// quantizer matrix (§7.12.3 step 1b): `q2 = Round2(q * matrix[i * tw + j],
1025    /// AOM_QM_BITS)`. `matrix` is [`quantizer_matrix`]'s slice for the block,
1026    /// or `None` when no matrix applies.
1027    #[must_use]
1028    pub fn dequantize_with_matrix(
1029        quant: &[i32],
1030        tx_size: TxSize,
1031        dc_quant: i64,
1032        ac_quant: i64,
1033        matrix: Option<&[u8]>,
1034        bit_depth: u8,
1035    ) -> Dequant {
1036        let tw = tx_size.width().min(32);
1037        let th = tx_size.height().min(32);
1038        let denom = tx_size.dq_denom();
1039        let mut values = [0_i64; 32 * 32];
1040        for (idx, out) in values.iter_mut().enumerate().take(tw * th) {
1041            let level = quant.get(idx).copied().unwrap_or(0);
1042            let q = if idx == 0 { dc_quant } else { ac_quant };
1043            let q = match matrix.and_then(|m| m.get(idx)) {
1044                Some(&weight) => (q * i64::from(weight) + (1 << (AOM_QM_BITS - 1))) >> AOM_QM_BITS,
1045                None => q,
1046            };
1047            let dq = i64::from(level) * q;
1048            let sign = if dq < 0 { -1 } else { 1 };
1049            let dq2 = sign * ((dq.abs() & 0xFF_FFFF) / denom);
1050            *out = clip_signed(dq2, 8 + u32::from(bit_depth));
1051        }
1052        Dequant {
1053            width: tw,
1054            height: th,
1055            values,
1056        }
1057    }
1058
1059    /// `Dc_Qlookup[(BitDepth-8)>>1][Clip3(0,255,b)]` (§7.12.2).
1060    #[must_use]
1061    pub fn dc_q(bit_depth: u8, b: i32) -> i64 {
1062        let row = usize::from(bit_depth.saturating_sub(8) >> 1).min(2);
1063        let col = b.clamp(0, 255) as usize;
1064        i64::from(DC_QLOOKUP[row][col])
1065    }
1066
1067    /// `Ac_Qlookup[(BitDepth-8)>>1][Clip3(0,255,b)]` (§7.12.2).
1068    #[must_use]
1069    pub fn ac_q(bit_depth: u8, b: i32) -> i64 {
1070        let row = usize::from(bit_depth.saturating_sub(8) >> 1).min(2);
1071        let col = b.clamp(0, 255) as usize;
1072        i64::from(AC_QLOOKUP[row][col])
1073    }
1074
1075    include!("quant_tables.rs");
1076    include!("qm_tables.rs");
1077
1078    /// `AOM_QM_BITS` (§3): the fixed-point precision of a matrix weight, in
1079    /// which 32 is unity.
1080    const AOM_QM_BITS: u32 = 5;
1081
1082    /// The quantizer matrix weights for a `tx_size` block, `Min(32, w) x
1083    /// Min(32, h)` of them row-major, at `level` for luma or chroma — or `None`
1084    /// at level 15, which means no matrix (`SegQMLevel`, §5.9.12).
1085    #[must_use]
1086    pub fn quantizer_matrix(level: u8, chroma: bool, tx_size: TxSize) -> Option<&'static [u8]> {
1087        let table = QUANTIZER_MATRIX.get(usize::from(level))?;
1088        let plane = table.get(usize::from(chroma))?;
1089        let start = usize::from(*QM_OFFSET.get(tx_size as usize)?);
1090        let len = tx_size.width().min(32) * tx_size.height().min(32);
1091        plane.get(start..start + len)
1092    }
1093
1094    #[cfg(test)]
1095    #[allow(
1096        clippy::unwrap_used,
1097        clippy::panic,
1098        reason = "tests operate on known-good values and assert shapes directly"
1099    )]
1100    mod dsp_tests {
1101        use super::*;
1102
1103        #[test]
1104        fn the_derived_forward_transform_inverts_the_decoders() {
1105            // Random residuals through forward, a fine quantizer, and the real
1106            // inverse come back within rounding, for every size and type the
1107            // encoder uses.
1108            let mut state = 0x1234_5678_u32;
1109            let types = [
1110                TxType::DctDct,
1111                TxType::AdstDct,
1112                TxType::DctAdst,
1113                TxType::AdstAdst,
1114                TxType::FlipadstDct,
1115            ];
1116            for size in [
1117                TxSize::Tx4x4,
1118                TxSize::Tx8x8,
1119                TxSize::Tx16x16,
1120                TxSize::Tx32x32,
1121                TxSize::Tx8x16,
1122                TxSize::Tx16x8,
1123                TxSize::Tx4x16,
1124            ] {
1125                for tx_type in types {
1126                    if size.sqr_up_idx() >= 3 && tx_type != TxType::DctDct {
1127                        continue;
1128                    }
1129                    let (w, h) = (size.width(), size.height());
1130                    let residual: Vec<i32> = (0..w * h)
1131                        .map(|_| {
1132                            state ^= state << 13;
1133                            state ^= state >> 17;
1134                            state ^= state << 5;
1135                            (state % 201) as i32 - 100
1136                        })
1137                        .collect();
1138                    let basis = ForwardBasis::new(size, tx_type);
1139                    let coeffs = basis.forward(&residual);
1140                    let denom = size.dq_denom();
1141                    let levels: Vec<i32> = coeffs
1142                        .iter()
1143                        .map(|&c| quantize(c, denom, size, 0.5))
1144                        .collect();
1145                    let dq = dequantize_with_matrix(&levels, size, denom, denom, None, 8);
1146                    let back = inverse_transform_2d(&dq, size, tx_type, false, 8);
1147                    let mut worst = 0;
1148                    for y in 0..h {
1149                        for x in 0..w {
1150                            let (yy, xx) = (
1151                                if tx_type.flip_ud() { h - 1 - y } else { y },
1152                                if tx_type.flip_lr() { w - 1 - x } else { x },
1153                            );
1154                            worst = worst.max((back.at(y, x) - residual[yy * w + xx]).abs());
1155                        }
1156                    }
1157                    assert!(worst <= 2, "{size:?} {tx_type:?}: off by {worst}");
1158                }
1159            }
1160        }
1161
1162        #[test]
1163        fn quantizer_matrix_lookup_follows_the_spec_table() {
1164            // Level 0 luma 4x4 opens the spec's table (§9.5.3).
1165            let m = quantizer_matrix(0, false, TxSize::Tx4x4).unwrap();
1166            assert_eq!(&m[..4], &[32, 43, 73, 97]);
1167            assert_eq!(m.len(), 16);
1168            // Sizes past 32 share the 32-capped matrix (Qm_Offset repeats).
1169            assert_eq!(
1170                quantizer_matrix(4, true, TxSize::Tx64x64),
1171                quantizer_matrix(4, true, TxSize::Tx32x32)
1172            );
1173            assert_eq!(
1174                quantizer_matrix(4, true, TxSize::Tx16x64).unwrap().len(),
1175                16 * 32
1176            );
1177            // Level 15 is "no matrix".
1178            assert_eq!(quantizer_matrix(15, false, TxSize::Tx8x8), None);
1179        }
1180
1181        #[test]
1182        fn a_flat_matrix_weight_leaves_the_quantizer_unchanged() {
1183            // 32 is unity at AOM_QM_BITS = 5; 48 is 1.5x, rounded.
1184            let quant = [3, -2, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0];
1185            let plain = dequantize(&quant, TxSize::Tx4x4, 40, 25, 8);
1186            let unity = dequantize_with_matrix(&quant, TxSize::Tx4x4, 40, 25, Some(&[32; 16]), 8);
1187            assert_eq!(plain.values, unity.values);
1188            let steeper = dequantize_with_matrix(&quant, TxSize::Tx4x4, 40, 25, Some(&[48; 16]), 8);
1189            assert_eq!(&steeper.values[..4], &[3 * 60, -2 * 38, 0, 38]);
1190        }
1191
1192        fn dequant_from(vals: &[(usize, i64)], tw: usize, th: usize) -> Dequant {
1193            let mut values = [0_i64; 32 * 32];
1194            for &(idx, v) in vals {
1195                values[idx] = v;
1196            }
1197            Dequant {
1198                width: tw,
1199                height: th,
1200                values,
1201            }
1202        }
1203
1204        #[test]
1205        fn cos_and_sin_hit_the_reference_points() {
1206            assert_eq!(cos128(0), 4096);
1207            assert_eq!(cos128(64), 0);
1208            assert_eq!(cos128(128), -4096);
1209            assert_eq!(sin128(64), 4096);
1210            assert_eq!(sin128(0), 0);
1211        }
1212
1213        #[test]
1214        fn brev_reverses_bits() {
1215            assert_eq!(brev(4, 1), 8);
1216            assert_eq!(brev(4, 0b0011), 0b1100);
1217            assert_eq!(brev(3, 0b001), 0b100);
1218        }
1219
1220        #[test]
1221        fn identity_identity_scales_a_dc_block() {
1222            // 8x8 IDTX: each pass multiplies by 2, then colShift = 4 halves twice.
1223            // A single DC term becomes (dc*2*2) >> 4 spread only at [0][0].
1224            let dequant = dequant_from(&[(0, 32)], 8, 8);
1225            let res = inverse_transform_2d(&dequant, TxSize::Tx8x8, TxType::Idtx, false, 8);
1226            assert_eq!(res.width, 8);
1227            assert_eq!(res.height, 8);
1228            // Row pass: T[0]=32*2=64, rowShift=1 -> 32. Col pass: 32*2=64,
1229            // colShift=4 -> 4. Only column 0, row 0 is non-zero.
1230            assert_eq!(res.at(0, 0), 4);
1231            assert_eq!(res.at(0, 1), 0);
1232            assert_eq!(res.at(1, 0), 0);
1233        }
1234
1235        #[test]
1236        fn dct_of_a_dc_only_block_is_flat() {
1237            // A DCT with only the DC coefficient set reconstructs a constant
1238            // block: every sample equal, no spatial variation.
1239            let dequant = dequant_from(&[(0, 512)], 8, 8);
1240            let res = inverse_transform_2d(&dequant, TxSize::Tx8x8, TxType::DctDct, false, 8);
1241            let first = res.at(0, 0);
1242            assert!(first != 0, "DC should reconstruct a non-zero level");
1243            for i in 0..8 {
1244                for j in 0..8 {
1245                    assert_eq!(res.at(i, j), first, "DCT DC block must be flat");
1246                }
1247            }
1248        }
1249
1250        #[test]
1251        fn adst_dc_block_is_not_flat_but_symmetric_is_valid() {
1252            // ADST is not flat for a DC input; just assert it runs and produces
1253            // a populated 4x4 block without panicking.
1254            let dequant = dequant_from(&[(0, 256)], 4, 4);
1255            let res = inverse_transform_2d(&dequant, TxSize::Tx4x4, TxType::AdstAdst, false, 8);
1256            assert_eq!(res.width, 4);
1257            assert_eq!(res.height, 4);
1258        }
1259
1260        #[test]
1261        fn lossless_dc_divides_the_dequantiser_back_out() {
1262            // Lossless is bit-exact because the qindex-0 dequant (times 4) and
1263            // the WHT's shift = 2 cancel: a DC level of 16 -> dequant 64 -> a
1264            // flat residual of 4, integrally, with no rounding loss.
1265            let dequant = dequant_from(&[(0, 64)], 4, 4);
1266            let res = inverse_transform_2d(&dequant, TxSize::Tx4x4, TxType::DctDct, true, 8);
1267            for i in 0..4 {
1268                for j in 0..4 {
1269                    assert_eq!(res.at(i, j), 4, "lossless DC must be flat and integral");
1270                }
1271            }
1272        }
1273
1274        #[test]
1275        fn dequantize_applies_dc_and_ac_steps() {
1276            // level 3 at DC with dc=10, level 2 at pos 1 with ac=5.
1277            let quant = [3_i32, 2, 0, 0];
1278            let dq = dequantize(&quant, TxSize::Tx4x4, 10, 5, 8);
1279            assert_eq!(dq.at(0, 0), 30);
1280            assert_eq!(dq.at(0, 1), 10);
1281        }
1282
1283        #[test]
1284        fn quant_lookups_hit_known_entries() {
1285            assert_eq!(dc_q(8, 0), 4);
1286            assert_eq!(ac_q(8, 0), 4);
1287            assert_eq!(dc_q(8, 255), 1336);
1288            assert_eq!(ac_q(8, 255), 1828);
1289            assert_eq!(dc_q(10, 0), 4);
1290            assert_eq!(dc_q(10, 255), 5347);
1291        }
1292    }
1293}
1294
1295#[cfg(test)]
1296#[allow(
1297    clippy::unwrap_used,
1298    clippy::indexing_slicing,
1299    clippy::panic,
1300    reason = "tests operate on known-good values and assert shapes directly"
1301)]
1302mod tests {
1303    use super::*;
1304
1305    /// Reconstruct a lossless 4x4 residual from raw coefficient levels, the way
1306    /// the tile driver does: qindex-0 dequant (`dc == ac == 4`) then the WHT.
1307    fn lossless_residual(quant: &[i32; 16]) -> Residual {
1308        let dq = dequantize(quant, TxSize::Tx4x4, 4, 4, 8);
1309        inverse_transform_2d(&dq, TxSize::Tx4x4, TxType::DctDct, true, 8)
1310    }
1311
1312    #[test]
1313    fn all_zero_coefficients_give_a_zero_residual() {
1314        let residual = lossless_residual(&[0; 16]);
1315        for i in 0..4 {
1316            for j in 0..4 {
1317                assert_eq!(residual.at(i, j), 0);
1318            }
1319        }
1320    }
1321
1322    #[test]
1323    fn a_dc_only_coefficient_spreads_evenly() {
1324        // A single DC level of 8 spreads to a flat block of 2.
1325        let mut quant = [0_i32; 16];
1326        quant[0] = 8;
1327        let residual = lossless_residual(&quant);
1328        for i in 0..4 {
1329            for j in 0..4 {
1330                assert_eq!(residual.at(i, j), 2, "DC residual should be flat");
1331            }
1332        }
1333    }
1334
1335    #[test]
1336    fn add_residual_clips_to_the_sample_range() {
1337        let pred = [[250_u16; 4]; 4];
1338        // A DC of 400 spreads to a flat +100; 250 + 100 clips to 255.
1339        let mut hi = [0_i32; 16];
1340        hi[0] = 400;
1341        let out = add_residual_4x4(&pred, &lossless_residual(&hi), 8);
1342        assert_eq!(out[0][0], 255);
1343        // A DC of -1200 spreads to a flat -300; 250 - 300 clips to 0.
1344        let mut lo = [0_i32; 16];
1345        lo[0] = -1200;
1346        let out = add_residual_4x4(&pred, &lossless_residual(&lo), 8);
1347        assert_eq!(out[1][1], 0);
1348    }
1349
1350    #[test]
1351    fn a_divisible_dc_reconstructs_integrally() {
1352        let mut quant = [0_i32; 16];
1353        quant[0] = 16;
1354        let residual = lossless_residual(&quant);
1355        assert_eq!(residual.at(0, 0), 4);
1356    }
1357
1358    #[test]
1359    fn flip_ud_places_the_residual_vertically_mirrored() {
1360        // With a uniform prediction, adding a residual under FLIPADST_DCT
1361        // (flipUD, no flipLR) must land Residual[i][j] at output row h-1-i,
1362        // i.e. the vertical mirror of the no-flip result — regardless of the
1363        // residual's values. Small residual + mid prediction avoids clipping.
1364        let pred = [128_u16; 16];
1365        let mut quant = [0_i32; 16];
1366        quant[1] = 8;
1367        quant[4] = -8;
1368        let res = lossless_residual(&quant);
1369        let no_flip = add_residual(&pred, &res, TxType::DctDct, 8);
1370        let flipped = add_residual(&pred, &res, TxType::FlipadstDct, 8);
1371        for i in 0..4 {
1372            for j in 0..4 {
1373                assert_eq!(flipped[(3 - i) * 4 + j], no_flip[i * 4 + j]);
1374            }
1375        }
1376    }
1377
1378    #[test]
1379    fn flip_lr_places_the_residual_horizontally_mirrored() {
1380        let pred = [128_u16; 16];
1381        let mut quant = [0_i32; 16];
1382        quant[1] = 8;
1383        quant[4] = -8;
1384        let res = lossless_residual(&quant);
1385        let no_flip = add_residual(&pred, &res, TxType::DctDct, 8);
1386        // DCT_FLIPADST is flipLR, no flipUD.
1387        let flipped = add_residual(&pred, &res, TxType::DctFlipadst, 8);
1388        for i in 0..4 {
1389            for j in 0..4 {
1390                assert_eq!(flipped[i * 4 + (3 - j)], no_flip[i * 4 + j]);
1391            }
1392        }
1393    }
1394}