1#[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#[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
82pub(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 fn round2(x: i64, n: u32) -> i64 {
102 if n == 0 { x } else { (x + (1 << (n - 1))) >> n }
103 }
104
105 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 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 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 fn sin128(angle: i32) -> i64 {
135 cos128(angle - 64)
136 }
137
138 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 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 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 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 #[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 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 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 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 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 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 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 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 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 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 #[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 #[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 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 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 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 #[must_use]
611 pub fn from_index(index: usize) -> TxSize {
612 ALL_TX_SIZES.get(index).copied().unwrap_or(TxSize::Tx4x4)
613 }
614
615 #[must_use]
617 pub fn log2_width(self) -> u32 {
618 TX_WIDTH_LOG2[self as usize]
619 }
620
621 #[must_use]
623 pub fn log2_height(self) -> u32 {
624 TX_HEIGHT_LOG2[self as usize]
625 }
626
627 #[must_use]
629 pub fn sqr_idx(self) -> u32 {
630 TX_SIZE_SQR[self as usize]
631 }
632
633 #[must_use]
635 pub fn sqr_up_idx(self) -> u32 {
636 TX_SIZE_SQR_UP[self as usize]
637 }
638
639 #[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 #[must_use]
649 pub fn adjusted_log2_width(self) -> u32 {
650 TX_WIDTH_LOG2[ADJUSTED_TX_SIZE[self as usize]]
651 }
652
653 #[must_use]
655 pub fn adjusted_height(self) -> usize {
656 1 << TX_HEIGHT_LOG2[ADJUSTED_TX_SIZE[self as usize]]
657 }
658
659 #[must_use]
661 pub fn adjusted_width(self) -> usize {
662 1 << self.adjusted_log2_width()
663 }
664
665 #[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 #[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 #[must_use]
682 pub fn width(self) -> usize {
683 1 << self.log2_width()
684 }
685
686 #[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 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 #[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 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 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 #[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 #[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 pub struct Residual {
784 pub width: usize,
786 pub height: usize,
788 values: [i32; 64 * 64],
789 }
790
791 impl Residual {
792 #[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 #[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 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 for value in residual.iter_mut().take(w * h) {
853 *value = clip_signed(*value, col_clamp);
854 }
855
856 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 #[derive(Debug, Clone)]
889 pub(crate) struct ForwardBasis {
890 w: usize,
891 h: usize,
892 row: Vec<f64>,
894 col: Vec<f64>,
896 row_norm: Vec<f64>,
898 col_norm: Vec<f64>,
899 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 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 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 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 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 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 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 #[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 #[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 #[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 #[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 const AOM_QM_BITS: u32 = 5;
1081
1082 #[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 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 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 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 assert_eq!(quantizer_matrix(15, false, TxSize::Tx8x8), None);
1179 }
1180
1181 #[test]
1182 fn a_flat_matrix_weight_leaves_the_quantizer_unchanged() {
1183 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 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 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 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 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 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 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 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 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 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 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 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 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}