Skip to main content

otf_pixels_codec_avif/av1/
coeff.rs

1//! Coefficient decode (spec §5.11.39, "coeffs") for every transform size.
2//!
3//! A transform block's residual is coded as a sparse level array `Quant[]`. The
4//! syntax reads, in order: `all_zero` (the whole block is zero), then the
5//! end-of-block position (`eob_pt` in one of seven size-dependent alphabets plus
6//! refinement bits), then the base level of each coefficient walking the scan
7//! *backwards* from the last, then a magnitude extension (`coeff_br` and a
8//! Golomb tail) for large levels, and finally the signs walking forwards. Every
9//! symbol is entropy-coded against a context-selected CDF, and getting the
10//! contexts exactly right is what keeps the arithmetic decoder in sync — a
11//! single wrong context desynchronises the whole tile.
12//!
13//! The contexts depend on the transform size (`txSzCtx`, the coded block's
14//! `bwl`) and the transform *class* — whether the transform is separable 2D or a
15//! 1D row/column identity transform, which reshapes the neighbour offsets and
16//! scan order. The lossless path is the `TX_4X4` / `DCT_DCT` (2D) corner of this
17//! and flows through the same entry point.
18
19use super::cdf;
20use super::symbol::{SymbolDecoder, SymbolEncoder};
21use super::transform::{TxSize, TxType};
22use super::transform_type::{
23    IntraTxSet, IntraTxTypeCdfs, chroma_tx_type, read_transform_type, write_transform_type,
24};
25use otf_pixels_core::{PixelsError, Result};
26
27include!("scan_tables.rs");
28
29/// Pick element `q` (0..=3) of a four-entry table by value, without indexing.
30fn pick4<T: Copy>(arr: [T; 4], q: usize) -> T {
31    let [a, b, c, d] = arr;
32    match q {
33        1 => b,
34        2 => c,
35        3 => d,
36        _ => a,
37    }
38}
39
40/// Borrow `slice[index]` as a mutable CDF row, or report a malformed stream
41/// rather than panic. Every call site derives its index from spec-bounded
42/// context maths, so this only fires on genuinely corrupt input.
43fn cdf_row<T>(slice: &mut [T], index: usize) -> Result<&mut T> {
44    slice.get_mut(index).ok_or_else(|| {
45        PixelsError::malformed("avif", "an AV1 coefficient CDF index ran out of range")
46    })
47}
48
49/// `NUM_BASE_LEVELS` (§3): base levels coded before the range extension.
50const NUM_BASE_LEVELS: i32 = 2;
51/// `COEFF_BASE_RANGE` (§3): the span covered by `coeff_br` before Golomb.
52const COEFF_BASE_RANGE: i32 = 12;
53/// `BR_CDF_SIZE` (§3): symbols in the `coeff_br` alphabet.
54const BR_CDF_SIZE: i32 = 4;
55/// `SIG_COEF_CONTEXTS` (§3).
56const SIG_COEF_CONTEXTS: usize = 42;
57/// `SIG_COEF_CONTEXTS_2D` (§3).
58const SIG_COEF_CONTEXTS_2D: i32 = 26;
59/// `SIG_COEF_CONTEXTS_EOB` (§3).
60const SIG_COEF_CONTEXTS_EOB: usize = 4;
61/// The largest scan a block can code (`TX_32X32`), the `Quant[]` capacity.
62const MAX_COEFFS: usize = 1024;
63
64/// Transform class (`get_tx_class`, §8.3.3): 0 = 2D, 1 = horizontal (1D row
65/// identity), 2 = vertical (1D column identity). The class selects the neighbour
66/// offsets and the scan order.
67fn tx_class(tx_type: TxType) -> usize {
68    match tx_type {
69        TxType::VDct | TxType::VAdst | TxType::VFlipadst => 2,
70        TxType::HDct | TxType::HAdst | TxType::HFlipadst => 1,
71        _ => 0,
72    }
73}
74
75/// `Sig_Ref_Diff_Offset[txClass][idx]` (§8.3.3): the neighbour offsets whose
76/// magnitudes drive the `coeff_base` context, as `(rowDelta, colDelta)`.
77const SIG_REF_DIFF_OFFSET: [[(i32, i32); 5]; 3] = [
78    [(0, 1), (1, 0), (1, 1), (0, 2), (2, 0)],
79    [(0, 1), (1, 0), (0, 2), (0, 3), (0, 4)],
80    [(0, 1), (1, 0), (2, 0), (3, 0), (4, 0)],
81];
82
83/// `Mag_Ref_Offset_With_Tx_Class[txClass][idx]` (§8.3.3): the neighbour offsets
84/// whose magnitudes drive the `coeff_br` context.
85const MAG_REF_OFFSET: [[(i32, i32); 3]; 3] = [
86    [(0, 1), (1, 0), (1, 1)],
87    [(0, 1), (1, 0), (0, 2)],
88    [(0, 1), (1, 0), (2, 0)],
89];
90
91/// `Coeff_Base_Pos_Ctx_Offset[Min(idx,2)]` (§8.3.3) for the 1D transform classes.
92const COEFF_BASE_POS_CTX_OFFSET: [i32; 3] = [
93    SIG_COEF_CONTEXTS_2D,
94    SIG_COEF_CONTEXTS_2D + 5,
95    SIG_COEF_CONTEXTS_2D + 10,
96];
97
98// `Coeff_Base_Ctx_Offset[txSz][Min(row,4)][Min(col,4)]` (§8.3.3). The 19 sizes
99// share six distinct patterns, named here to make the shape auditable against
100// the spec: `SQR` for the square sizes, `TALL`/`WIDE` for the 2:1 and 4:1
101// rectangles, and the small 4x4 / 4x8 / 8x4 variants that zero their last
102// row or column.
103const CBO_4X4: [[i32; 5]; 5] = [
104    [0, 1, 6, 6, 0],
105    [1, 6, 6, 21, 0],
106    [6, 6, 21, 21, 0],
107    [6, 21, 21, 21, 0],
108    [0, 0, 0, 0, 0],
109];
110const CBO_SQR: [[i32; 5]; 5] = [
111    [0, 1, 6, 6, 21],
112    [1, 6, 6, 21, 21],
113    [6, 6, 21, 21, 21],
114    [6, 21, 21, 21, 21],
115    [21, 21, 21, 21, 21],
116];
117const CBO_NARROW: [[i32; 5]; 5] = [
118    [0, 11, 11, 11, 0],
119    [11, 11, 11, 11, 0],
120    [6, 6, 21, 21, 0],
121    [6, 21, 21, 21, 0],
122    [21, 21, 21, 21, 0],
123];
124const CBO_SHORT: [[i32; 5]; 5] = [
125    [0, 16, 6, 6, 21],
126    [16, 16, 6, 21, 21],
127    [16, 16, 21, 21, 21],
128    [16, 16, 21, 21, 21],
129    [0, 0, 0, 0, 0],
130];
131const CBO_TALL: [[i32; 5]; 5] = [
132    [0, 11, 11, 11, 11],
133    [11, 11, 11, 11, 11],
134    [6, 6, 21, 21, 21],
135    [6, 21, 21, 21, 21],
136    [21, 21, 21, 21, 21],
137];
138const CBO_WIDE: [[i32; 5]; 5] = [
139    [0, 16, 6, 6, 21],
140    [16, 16, 6, 21, 21],
141    [16, 16, 21, 21, 21],
142    [16, 16, 21, 21, 21],
143    [16, 16, 21, 21, 21],
144];
145const COEFF_BASE_CTX_OFFSET: [[[i32; 5]; 5]; 19] = [
146    CBO_4X4, CBO_SQR, CBO_SQR, CBO_SQR, CBO_SQR, CBO_NARROW, CBO_SHORT, CBO_TALL, CBO_WIDE,
147    CBO_TALL, CBO_WIDE, CBO_TALL, CBO_WIDE, CBO_NARROW, CBO_SHORT, CBO_TALL, CBO_WIDE, CBO_TALL,
148    CBO_WIDE,
149];
150
151/// The mutable coefficient CDFs for one tile, cloned from the defaults for the
152/// frame's quantiser context. The spec's `Tile*Cdf` are the frame defaults
153/// pre-indexed by the quantiser context (`get_qctx`), then adapted per symbol as
154/// the tile decodes.
155pub struct CoeffCdfs {
156    txb_skip: [[[u16; 3]; 13]; 5],
157    eob_pt_16: [[[u16; 6]; 2]; 2],
158    eob_pt_32: [[[u16; 7]; 2]; 2],
159    eob_pt_64: [[[u16; 8]; 2]; 2],
160    eob_pt_128: [[[u16; 9]; 2]; 2],
161    eob_pt_256: [[[u16; 10]; 2]; 2],
162    eob_pt_512: [[u16; 11]; 2],
163    eob_pt_1024: [[u16; 12]; 2],
164    eob_extra: [[[[u16; 3]; 9]; 2]; 5],
165    coeff_base_eob: [[[[u16; 4]; 4]; 2]; 5],
166    coeff_base: [[[[u16; 5]; 42]; 2]; 5],
167    coeff_br: [[[[u16; 5]; 21]; 2]; 5],
168    dc_sign: [[[u16; 3]; 3]; 2],
169}
170
171impl CoeffCdfs {
172    /// Clone the defaults for quantiser context `qctx` (0 for lossless).
173    #[must_use]
174    pub fn new(qctx: usize) -> Self {
175        let q = qctx.min(3);
176        Self {
177            txb_skip: pick4(cdf::DEFAULT_TXB_SKIP_CDF, q),
178            eob_pt_16: pick4(cdf::DEFAULT_EOB_PT_16_CDF, q),
179            eob_pt_32: pick4(cdf::DEFAULT_EOB_PT_32_CDF, q),
180            eob_pt_64: pick4(cdf::DEFAULT_EOB_PT_64_CDF, q),
181            eob_pt_128: pick4(cdf::DEFAULT_EOB_PT_128_CDF, q),
182            eob_pt_256: pick4(cdf::DEFAULT_EOB_PT_256_CDF, q),
183            eob_pt_512: pick4(cdf::DEFAULT_EOB_PT_512_CDF, q),
184            eob_pt_1024: pick4(cdf::DEFAULT_EOB_PT_1024_CDF, q),
185            eob_extra: pick4(cdf::DEFAULT_EOB_EXTRA_CDF, q),
186            coeff_base_eob: pick4(cdf::DEFAULT_COEFF_BASE_EOB_CDF, q),
187            coeff_base: pick4(cdf::DEFAULT_COEFF_BASE_CDF, q),
188            coeff_br: pick4(cdf::DEFAULT_COEFF_BR_CDF, q),
189            dc_sign: pick4(cdf::DEFAULT_DC_SIGN_CDF, q),
190        }
191    }
192}
193
194/// The result of decoding one transform block's coefficients.
195pub struct CoeffBlock {
196    /// `Quant[]` in the coded block's raster order: signed dequantiser input
197    /// levels. Only the first `Tx_Width * Tx_Height` (of the adjusted size) are
198    /// meaningful; the tail stays zero.
199    pub quant: [i32; MAX_COEFFS],
200    /// The end-of-block position: the count of leading scan coefficients.
201    pub eob: usize,
202    /// `culLevel`, clamped to 63: the neighbour level context this block leaves.
203    pub cul_level: u8,
204    /// `dcCategory`: 0 none, 1 negative DC, 2 positive DC.
205    pub dc_category: u8,
206    /// The resolved `PlaneTxType` (`compute_tx_type`, §5.11.40): `DCT_DCT` for a
207    /// skipped or lossless block, otherwise the type read/derived here. The
208    /// dequantiser and inverse transform need it, so it travels with the block.
209    pub tx_type: TxType,
210}
211
212/// How a coded block's `PlaneTxType` is resolved once `all_zero` shows the block
213/// carries coefficients (`compute_tx_type` / `transform_type`, §5.11.40–47).
214///
215/// Luma reads the `intra_tx_type` symbol at this point — its spec position,
216/// after `all_zero` and before `eob_pt`; chroma derives its type from the
217/// prediction mode without consuming a symbol. A zero quantiser (the lossless
218/// path) or a `DCT_DCT`-only set forces `DCT_DCT`.
219pub struct TxTypeCtx<'a> {
220    /// The intra transform set for this size (`get_tx_set`).
221    pub set: IntraTxSet,
222    /// The adapting `intra_tx_type` CDFs; only luma reads them.
223    pub intra_cdfs: &'a mut IntraTxTypeCdfs,
224    /// `intraDir` for the luma symbol's context.
225    pub intra_dir: usize,
226    /// `UVMode` for the chroma (plane > 0) derivation.
227    pub uv_mode: usize,
228    /// Whether the segment's quantiser index (`get_qindex(1, segment_id)`) is
229    /// non-zero; only then is a luma `intra_tx_type` symbol coded.
230    pub qindex_positive: bool,
231    /// The block's `Lossless`: every plane is then `DCT_DCT` (the WHT path).
232    pub lossless: bool,
233}
234
235/// The scan order (`get_scan`, §5.11.41) for a transform size and type.
236fn get_scan(tx_size: TxSize, tx_type: TxType) -> &'static [u16] {
237    match tx_size {
238        TxSize::Tx16x64 => return &DEFAULT_SCAN_16X32,
239        TxSize::Tx64x16 => return &DEFAULT_SCAN_32X16,
240        _ => {}
241    }
242    if tx_size.sqr_up_idx() == 4 {
243        // Tx_Size_Sqr_Up == TX_64X64: coded as a 32x32 scan.
244        return &DEFAULT_SCAN_32X32;
245    }
246    match tx_type {
247        TxType::Idtx => default_scan(tx_size),
248        TxType::VDct | TxType::VAdst | TxType::VFlipadst => mrow_scan(tx_size),
249        TxType::HDct | TxType::HAdst | TxType::HFlipadst => mcol_scan(tx_size),
250        _ => default_scan(tx_size),
251    }
252}
253
254/// `get_default_scan(txSz)` (§5.11.41).
255fn default_scan(tx_size: TxSize) -> &'static [u16] {
256    match tx_size {
257        TxSize::Tx4x4 => &DEFAULT_SCAN_4X4,
258        TxSize::Tx4x8 => &DEFAULT_SCAN_4X8,
259        TxSize::Tx8x4 => &DEFAULT_SCAN_8X4,
260        TxSize::Tx8x8 => &DEFAULT_SCAN_8X8,
261        TxSize::Tx8x16 => &DEFAULT_SCAN_8X16,
262        TxSize::Tx16x8 => &DEFAULT_SCAN_16X8,
263        TxSize::Tx16x16 => &DEFAULT_SCAN_16X16,
264        TxSize::Tx16x32 => &DEFAULT_SCAN_16X32,
265        TxSize::Tx32x16 => &DEFAULT_SCAN_32X16,
266        TxSize::Tx4x16 => &DEFAULT_SCAN_4X16,
267        TxSize::Tx16x4 => &DEFAULT_SCAN_16X4,
268        TxSize::Tx8x32 => &DEFAULT_SCAN_8X32,
269        TxSize::Tx32x8 => &DEFAULT_SCAN_32X8,
270        _ => &DEFAULT_SCAN_32X32,
271    }
272}
273
274/// `get_mrow_scan(txSz)` (§5.11.41), used by the `V_*` (row-identity) types.
275fn mrow_scan(tx_size: TxSize) -> &'static [u16] {
276    match tx_size {
277        TxSize::Tx4x4 => &MROW_SCAN_4X4,
278        TxSize::Tx4x8 => &MROW_SCAN_4X8,
279        TxSize::Tx8x4 => &MROW_SCAN_8X4,
280        TxSize::Tx8x8 => &MROW_SCAN_8X8,
281        TxSize::Tx8x16 => &MROW_SCAN_8X16,
282        TxSize::Tx16x8 => &MROW_SCAN_16X8,
283        TxSize::Tx16x16 => &MROW_SCAN_16X16,
284        TxSize::Tx4x16 => &MROW_SCAN_4X16,
285        _ => &MROW_SCAN_16X4,
286    }
287}
288
289/// `get_mcol_scan(txSz)` (§5.11.41), used by the `H_*` (column-identity) types.
290fn mcol_scan(tx_size: TxSize) -> &'static [u16] {
291    match tx_size {
292        TxSize::Tx4x4 => &MCOL_SCAN_4X4,
293        TxSize::Tx4x8 => &MCOL_SCAN_4X8,
294        TxSize::Tx8x4 => &MCOL_SCAN_8X4,
295        TxSize::Tx8x8 => &MCOL_SCAN_8X8,
296        TxSize::Tx8x16 => &MCOL_SCAN_8X16,
297        TxSize::Tx16x8 => &MCOL_SCAN_16X8,
298        TxSize::Tx16x16 => &MCOL_SCAN_16X16,
299        TxSize::Tx4x16 => &MCOL_SCAN_4X16,
300        _ => &MCOL_SCAN_16X4,
301    }
302}
303
304/// Decode the coefficients of one transform block (`coeffs`, §5.11.39).
305///
306/// `ptype` is 0 for luma and 1 for chroma. `all_zero_ctx` and `dc_sign_ctx` are
307/// the neighbour-derived contexts the tile driver computes. `tx` resolves the
308/// `PlaneTxType`: for a coded luma block the `intra_tx_type` symbol is read here,
309/// at its spec position between `all_zero` and `eob_pt`. The returned `Quant[]`
310/// feeds the dequantiser and inverse transform, and the resolved type rides
311/// along on the block.
312///
313/// # Errors
314///
315/// Propagates any error from the arithmetic decoder (a stream that ends early)
316/// or a corrupt context index.
317pub fn decode_coeffs(
318    dec: &mut SymbolDecoder<'_>,
319    cdfs: &mut CoeffCdfs,
320    tx_size: TxSize,
321    tx: TxTypeCtx<'_>,
322    ptype: usize,
323    all_zero_ctx: usize,
324    dc_sign_ctx: usize,
325) -> Result<CoeffBlock> {
326    let mut quant = [0_i32; MAX_COEFFS];
327    let pt = ptype.min(1);
328    let tx_ctx = tx_size.tx_size_ctx();
329
330    // all_zero (txb_skip): the whole block codes as zero.
331    let skip_cdf = cdf_row(cdf_row(&mut cdfs.txb_skip, tx_ctx)?, all_zero_ctx)?;
332    let all_zero = dec.read_symbol(skip_cdf)? != 0;
333    if all_zero {
334        return Ok(CoeffBlock {
335            quant,
336            eob: 0,
337            cul_level: 0,
338            dc_category: 0,
339            tx_type: TxType::DctDct,
340        });
341    }
342
343    // transform_type (§5.11.40/47): resolve the PlaneTxType now — luma reads the
344    // intra_tx_type symbol at this point, chroma derives from its prediction
345    // mode without a symbol. Lossless (qindex 0) forces DCT_DCT either way.
346    let tx_type = if ptype == 0 {
347        read_transform_type(
348            dec,
349            tx.intra_cdfs,
350            tx.set,
351            tx_size,
352            tx.intra_dir,
353            tx.qindex_positive,
354        )?
355    } else if !tx.lossless {
356        chroma_tx_type(tx.uv_mode, tx.set)
357    } else {
358        TxType::DctDct
359    };
360    let cls = tx_class(tx_type);
361
362    let scan = get_scan(tx_size, tx_type);
363    let eob_ctx = usize::from(cls != 0);
364
365    // eob_pt: the end-of-block bucket, in a size-dependent alphabet.
366    let eob_pt = match tx_size.eob_multisize() {
367        0 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_16, pt)?, eob_ctx)?)?,
368        1 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_32, pt)?, eob_ctx)?)?,
369        2 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_64, pt)?, eob_ctx)?)?,
370        3 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_128, pt)?, eob_ctx)?)?,
371        4 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_256, pt)?, eob_ctx)?)?,
372        5 => dec.read_symbol(cdf_row(&mut cdfs.eob_pt_512, pt)?)?,
373        _ => dec.read_symbol(cdf_row(&mut cdfs.eob_pt_1024, pt)?)?,
374    } + 1;
375
376    let mut eob = if eob_pt < 2 {
377        eob_pt
378    } else {
379        (1 << (eob_pt - 2)) + 1
380    };
381
382    // eob_extra plus the raw extra bits refine eob within its bucket.
383    if let Some(eob_shift) = eob_pt.checked_sub(3) {
384        let extra_cdf = cdf_row(
385            cdf_row(cdf_row(&mut cdfs.eob_extra, tx_ctx)?, pt)?,
386            eob_pt - 3,
387        )?;
388        if dec.read_symbol(extra_cdf)? != 0 {
389            eob += 1 << eob_shift;
390        }
391        for i in 1..eob_pt.saturating_sub(2) {
392            let shift = eob_pt.saturating_sub(2) - 1 - i;
393            if dec.read_bool()? {
394                eob += 1 << shift;
395            }
396        }
397    }
398
399    eob = eob.min(scan.len());
400
401    // Base levels, walking the scan backwards from the last coefficient.
402    for c in (0..eob).rev() {
403        let pos = scan.get(c).map_or(0, |&p| usize::from(p));
404        let mut level;
405        if c == eob - 1 {
406            let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, true) + SIG_COEF_CONTEXTS_EOB
407                - SIG_COEF_CONTEXTS;
408            let cdf_ref = cdf_row(
409                cdf_row(cdf_row(&mut cdfs.coeff_base_eob, tx_ctx)?, pt)?,
410                ctx,
411            )?;
412            level = dec.read_symbol(cdf_ref)? as i32 + 1;
413        } else {
414            let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, false);
415            let cdf_ref = cdf_row(cdf_row(cdf_row(&mut cdfs.coeff_base, tx_ctx)?, pt)?, ctx)?;
416            level = dec.read_symbol(cdf_ref)? as i32;
417        }
418        if level > NUM_BASE_LEVELS {
419            let br_ctx = coeff_br_ctx(tx_size, cls, &quant, pos);
420            let br_bucket = tx_ctx.min(3);
421            for _ in 0..(COEFF_BASE_RANGE / (BR_CDF_SIZE - 1)) {
422                let cdf_ref = cdf_row(
423                    cdf_row(cdf_row(&mut cdfs.coeff_br, br_bucket)?, pt)?,
424                    br_ctx,
425                )?;
426                let coeff_br = dec.read_symbol(cdf_ref)? as i32;
427                level += coeff_br;
428                if coeff_br < BR_CDF_SIZE - 1 {
429                    break;
430                }
431            }
432        }
433        if let Some(slot) = quant.get_mut(pos) {
434            *slot = level;
435        }
436    }
437
438    // Signs and the Golomb magnitude tail, walking the scan forwards.
439    let mut cul_level: i32 = 0;
440    let mut dc_category = 0_u8;
441    for c in 0..eob {
442        let pos = scan.get(c).map_or(0, |&p| usize::from(p));
443        let level = quant.get(pos).copied().unwrap_or(0);
444        let sign = if level != 0 {
445            if c == 0 {
446                let cdf_ref = cdf_row(cdf_row(&mut cdfs.dc_sign, pt)?, dc_sign_ctx)?;
447                dec.read_symbol(cdf_ref)? != 0
448            } else {
449                dec.read_bool()?
450            }
451        } else {
452            false
453        };
454        let mut magnitude = level;
455        if magnitude > NUM_BASE_LEVELS + COEFF_BASE_RANGE {
456            magnitude = read_golomb(dec)? + COEFF_BASE_RANGE + NUM_BASE_LEVELS;
457        }
458        if pos == 0 && magnitude > 0 {
459            dc_category = if sign { 1 } else { 2 };
460        }
461        magnitude &= 0xF_FFFF;
462        cul_level += magnitude;
463        if let Some(slot) = quant.get_mut(pos) {
464            *slot = if sign { -magnitude } else { magnitude };
465        }
466    }
467
468    Ok(CoeffBlock {
469        quant,
470        eob,
471        cul_level: cul_level.min(63) as u8,
472        dc_category,
473        tx_type,
474    })
475}
476
477/// Write one transform block's coefficients where [`decode_coeffs`] would
478/// read them, and return the block exactly as the decoder will see it.
479///
480/// `levels` are the signed quantized levels in the block's raster order (the
481/// layout of [`CoeffBlock::quant`]); `tx_type` is what the encoder transformed
482/// with, which must be the type the decoder will resolve.
483///
484/// # Errors
485///
486/// Rejects an out-of-range context.
487#[allow(clippy::too_many_arguments, reason = "the coefficient syntax's inputs")]
488pub(crate) fn encode_coeffs(
489    enc: &mut SymbolEncoder,
490    cdfs: &mut CoeffCdfs,
491    tx_size: TxSize,
492    tx: TxTypeCtx<'_>,
493    tx_type: TxType,
494    ptype: usize,
495    all_zero_ctx: usize,
496    dc_sign_ctx: usize,
497    levels: &[i32],
498) -> Result<CoeffBlock> {
499    let pt = ptype.min(1);
500    let tx_ctx = tx_size.tx_size_ctx();
501    let cls = tx_class(tx_type);
502    let scan = get_scan(tx_size, tx_type);
503    let level_at = |c: usize| {
504        levels
505            .get(scan.get(c).map_or(0, |&p| usize::from(p)))
506            .copied()
507            .unwrap_or(0)
508    };
509    let eob = (0..scan.len())
510        .rev()
511        .find(|&c| level_at(c) != 0)
512        .map_or(0, |c| c + 1);
513
514    let skip_cdf = cdf_row(cdf_row(&mut cdfs.txb_skip, tx_ctx)?, all_zero_ctx)?;
515    enc.write_symbol(skip_cdf, usize::from(eob == 0));
516    let mut quant = [0_i32; MAX_COEFFS];
517    if eob == 0 {
518        return Ok(CoeffBlock {
519            quant,
520            eob: 0,
521            cul_level: 0,
522            dc_category: 0,
523            tx_type: TxType::DctDct,
524        });
525    }
526    if ptype == 0 {
527        write_transform_type(
528            enc,
529            tx.intra_cdfs,
530            tx.set,
531            tx_size,
532            tx.intra_dir,
533            tx.qindex_positive,
534            tx_type,
535        )?;
536    }
537
538    // eob_pt and its refinement bits.
539    let eob_pt = if eob <= 2 {
540        eob
541    } else {
542        2 + (usize::BITS - 1 - (eob - 1).leading_zeros()) as usize
543    };
544    let eob_ctx = usize::from(cls != 0);
545    let symbol = eob_pt - 1;
546    match tx_size.eob_multisize() {
547        0 => enc.write_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_16, pt)?, eob_ctx)?, symbol),
548        1 => enc.write_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_32, pt)?, eob_ctx)?, symbol),
549        2 => enc.write_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_64, pt)?, eob_ctx)?, symbol),
550        3 => enc.write_symbol(
551            cdf_row(cdf_row(&mut cdfs.eob_pt_128, pt)?, eob_ctx)?,
552            symbol,
553        ),
554        4 => enc.write_symbol(
555            cdf_row(cdf_row(&mut cdfs.eob_pt_256, pt)?, eob_ctx)?,
556            symbol,
557        ),
558        5 => enc.write_symbol(cdf_row(&mut cdfs.eob_pt_512, pt)?, symbol),
559        _ => enc.write_symbol(cdf_row(&mut cdfs.eob_pt_1024, pt)?, symbol),
560    }
561    if eob_pt >= 3 {
562        let offset = eob - ((1 << (eob_pt - 2)) + 1);
563        let extra_cdf = cdf_row(
564            cdf_row(cdf_row(&mut cdfs.eob_extra, tx_ctx)?, pt)?,
565            eob_pt - 3,
566        )?;
567        enc.write_symbol(extra_cdf, (offset >> (eob_pt - 3)) & 1);
568        for i in 1..eob_pt - 2 {
569            let shift = eob_pt - 2 - 1 - i;
570            enc.write_bool((offset >> shift) & 1 == 1);
571        }
572    }
573
574    // Base levels and ranges, backwards, filling `quant` as the decoder does.
575    let cap = NUM_BASE_LEVELS + COEFF_BASE_RANGE + 1;
576    for c in (0..eob).rev() {
577        let pos = scan.get(c).map_or(0, |&p| usize::from(p));
578        let magnitude = level_at(c).abs().min(cap);
579        if c == eob - 1 {
580            let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, true) + SIG_COEF_CONTEXTS_EOB
581                - SIG_COEF_CONTEXTS;
582            let cdf_ref = cdf_row(
583                cdf_row(cdf_row(&mut cdfs.coeff_base_eob, tx_ctx)?, pt)?,
584                ctx,
585            )?;
586            enc.write_symbol(cdf_ref, (magnitude.min(3) - 1) as usize);
587        } else {
588            let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, false);
589            let cdf_ref = cdf_row(cdf_row(cdf_row(&mut cdfs.coeff_base, tx_ctx)?, pt)?, ctx)?;
590            enc.write_symbol(cdf_ref, magnitude.min(3) as usize);
591        }
592        if magnitude > NUM_BASE_LEVELS {
593            let br_ctx = coeff_br_ctx(tx_size, cls, &quant, pos);
594            let br_bucket = tx_ctx.min(3);
595            let mut remaining = magnitude - NUM_BASE_LEVELS - 1;
596            for _ in 0..(COEFF_BASE_RANGE / (BR_CDF_SIZE - 1)) {
597                let k = remaining.min(BR_CDF_SIZE - 1);
598                let cdf_ref = cdf_row(
599                    cdf_row(cdf_row(&mut cdfs.coeff_br, br_bucket)?, pt)?,
600                    br_ctx,
601                )?;
602                enc.write_symbol(cdf_ref, k as usize);
603                remaining -= k;
604                if k < BR_CDF_SIZE - 1 {
605                    break;
606                }
607            }
608        }
609        if let Some(slot) = quant.get_mut(pos) {
610            *slot = magnitude;
611        }
612    }
613
614    // Signs and Golomb tails, forwards.
615    let mut cul_level = 0_i32;
616    let mut dc_category = 0_u8;
617    for c in 0..eob {
618        let pos = scan.get(c).map_or(0, |&p| usize::from(p));
619        let value = level_at(c);
620        let magnitude = value.abs();
621        if magnitude != 0 {
622            if c == 0 {
623                let cdf_ref = cdf_row(cdf_row(&mut cdfs.dc_sign, pt)?, dc_sign_ctx)?;
624                enc.write_symbol(cdf_ref, usize::from(value < 0));
625            } else {
626                enc.write_bool(value < 0);
627            }
628        }
629        if magnitude > NUM_BASE_LEVELS + COEFF_BASE_RANGE {
630            write_golomb(enc, (magnitude - NUM_BASE_LEVELS - COEFF_BASE_RANGE) as u32);
631        }
632        if pos == 0 && magnitude > 0 {
633            dc_category = if value < 0 { 1 } else { 2 };
634        }
635        let magnitude = magnitude & 0xF_FFFF;
636        cul_level += magnitude;
637        if let Some(slot) = quant.get_mut(pos) {
638            *slot = if value < 0 { -magnitude } else { magnitude };
639        }
640    }
641    Ok(CoeffBlock {
642        quant,
643        eob,
644        cul_level: cul_level.min(63) as u8,
645        dc_category,
646        tx_type,
647    })
648}
649
650/// The inverse of `read_golomb`: `x >= 1` as its bit length in zeros, then
651/// its bits from the top.
652fn write_golomb(enc: &mut SymbolEncoder, x: u32) {
653    let length = 32 - x.leading_zeros();
654    for _ in 1..length {
655        enc.write_bool(false);
656    }
657    enc.write_bool(true);
658    for i in (0..length - 1).rev() {
659        enc.write_bool((x >> i) & 1 == 1);
660    }
661}
662
663/// `get_coeff_base_ctx` (§8.3.3). `is_eob` selects the four end-of-block
664/// contexts; otherwise the magnitude of already-decoded scan neighbours and the
665/// coefficient position pick the context.
666fn coeff_base_ctx(
667    tx_size: TxSize,
668    cls: usize,
669    quant: &[i32; MAX_COEFFS],
670    pos: usize,
671    c: usize,
672    is_eob: bool,
673) -> usize {
674    let bwl = tx_size.adjusted_log2_width();
675    let width = tx_size.adjusted_width() as i32;
676    let height = tx_size.adjusted_height() as i32;
677    if is_eob {
678        let area = (height as usize) << bwl;
679        return if c == 0 {
680            SIG_COEF_CONTEXTS - 4
681        } else if c <= area / 8 {
682            SIG_COEF_CONTEXTS - 3
683        } else if c <= area / 4 {
684            SIG_COEF_CONTEXTS - 2
685        } else {
686            SIG_COEF_CONTEXTS - 1
687        };
688    }
689    let row = (pos >> bwl) as i32;
690    let col = (pos - ((row as usize) << bwl)) as i32;
691    let mut mag = 0;
692    for &(d_row, d_col) in offsets_sig(cls) {
693        let ref_row = row + d_row;
694        let ref_col = col + d_col;
695        if ref_row >= 0 && ref_col >= 0 && ref_row < height && ref_col < width {
696            let ref_pos = ((ref_row as usize) << bwl) + ref_col as usize;
697            mag += quant.get(ref_pos).copied().unwrap_or(0).abs().min(3);
698        }
699    }
700    let ctx = ((mag + 1) >> 1).min(4);
701    if cls == 0 {
702        if row == 0 && col == 0 {
703            return 0;
704        }
705        let offset = COEFF_BASE_CTX_OFFSET
706            .get(tx_size as usize)
707            .and_then(|t| t.get(row.min(4) as usize))
708            .and_then(|r| r.get(col.min(4) as usize))
709            .copied()
710            .unwrap_or(0);
711        return (ctx + offset) as usize;
712    }
713    let idx = if cls == 2 { row } else { col };
714    let offset = pick3(COEFF_BASE_POS_CTX_OFFSET, idx.min(2) as usize);
715    (ctx + offset) as usize
716}
717
718/// `coeff_br` context (§8.3.3).
719fn coeff_br_ctx(tx_size: TxSize, cls: usize, quant: &[i32; MAX_COEFFS], pos: usize) -> usize {
720    let bwl = tx_size.adjusted_log2_width();
721    let txw = tx_size.adjusted_width();
722    let txh = tx_size.adjusted_height() as i32;
723    let row = (pos >> bwl) as i32;
724    let col = (pos - ((row as usize) << bwl)) as i32;
725    let mut mag = 0;
726    for &(d_row, d_col) in offsets_mag(cls) {
727        let ref_row = row + d_row;
728        let ref_col = col + d_col;
729        if ref_row >= 0 && ref_col >= 0 && ref_row < txh && ref_col < (1 << bwl) {
730            let ref_pos = ref_row as usize * txw + ref_col as usize;
731            mag += quant
732                .get(ref_pos)
733                .copied()
734                .unwrap_or(0)
735                .min(COEFF_BASE_RANGE + NUM_BASE_LEVELS + 1);
736        }
737    }
738    let mag = ((mag + 1) >> 1).min(6);
739    let ctx = if pos == 0 {
740        mag
741    } else if cls == 0 {
742        if row < 2 && col < 2 {
743            mag + 7
744        } else {
745            mag + 14
746        }
747    } else if cls == 1 {
748        if col == 0 { mag + 7 } else { mag + 14 }
749    } else if row == 0 {
750        mag + 7
751    } else {
752        mag + 14
753    };
754    ctx as usize
755}
756
757/// `Sig_Ref_Diff_Offset[cls]` without indexing.
758fn offsets_sig(cls: usize) -> &'static [(i32, i32); 5] {
759    let [two_d, horiz, vert] = &SIG_REF_DIFF_OFFSET;
760    match cls {
761        1 => horiz,
762        2 => vert,
763        _ => two_d,
764    }
765}
766
767/// `Mag_Ref_Offset_With_Tx_Class[cls]` without indexing.
768fn offsets_mag(cls: usize) -> &'static [(i32, i32); 3] {
769    let [two_d, horiz, vert] = &MAG_REF_OFFSET;
770    match cls {
771        1 => horiz,
772        2 => vert,
773        _ => two_d,
774    }
775}
776
777/// Pick element `q` (0..=2) of a three-entry table by value, without indexing.
778fn pick3<T: Copy>(arr: [T; 3], q: usize) -> T {
779    let [a, b, c] = arr;
780    match q {
781        1 => b,
782        2 => c,
783        _ => a,
784    }
785}
786
787/// Read an exp-Golomb coded magnitude tail (`golomb`, §5.11.39).
788fn read_golomb(dec: &mut SymbolDecoder<'_>) -> Result<i32> {
789    let mut length = 0_i32;
790    loop {
791        length += 1;
792        if dec.read_bool()? {
793            break;
794        }
795        if length > 20 {
796            break;
797        }
798    }
799    let mut x = 1_i32;
800    for _ in 0..length.saturating_sub(1) {
801        x = (x << 1) | i32::from(dec.read_bool()?);
802    }
803    Ok(x)
804}
805
806#[cfg(test)]
807#[allow(
808    clippy::unwrap_used,
809    clippy::indexing_slicing,
810    clippy::panic,
811    reason = "tests operate on known-good values and assert shapes directly"
812)]
813mod tests {
814    use super::*;
815
816    #[test]
817    fn written_coefficients_decode_to_the_same_block() {
818        use super::super::transform_type::intra_tx_set;
819        let mut state = 0x9e37_79b9_u32;
820        let mut next = move || {
821            state ^= state << 13;
822            state ^= state >> 17;
823            state ^= state << 5;
824            state
825        };
826        let sizes = [
827            TxSize::Tx4x4,
828            TxSize::Tx8x8,
829            TxSize::Tx16x16,
830            TxSize::Tx32x32,
831            TxSize::Tx8x16,
832            TxSize::Tx16x4,
833        ];
834        for round in 0..60 {
835            let size = sizes[round % sizes.len()];
836            let ptype = usize::from(round % 3 == 2);
837            let set = intra_tx_set(size, false);
838            let uv_mode = 1; // V_PRED: chroma derives ADST_DCT where the set allows
839            let tx_type = if ptype == 0 {
840                TxType::DctDct
841            } else {
842                chroma_tx_type(uv_mode, set)
843            };
844            let (w, h) = (size.adjusted_width(), size.adjusted_height());
845            // Sparse levels, a few of them huge (Golomb tails), sometimes none.
846            let density = next() % 4;
847            let levels: Vec<i32> = (0..w * h)
848                .map(|_| {
849                    let r = next();
850                    if density == 0 || r % (2 + density * 3) != 0 {
851                        0
852                    } else {
853                        let m = match r % 7 {
854                            0 => 40 + (r >> 20) as i32 % 3000,
855                            1 => 3 + (r >> 8) as i32 % 12,
856                            _ => 1 + (r >> 8) as i32 % 2,
857                        };
858                        if r & 1 == 0 { m } else { -m }
859                    }
860                })
861                .collect();
862            let (az, ds) = ((next() % 13) as usize, (next() % 3) as usize);
863            let mut enc = SymbolEncoder::new(false);
864            let mut cdfs = CoeffCdfs::new(2);
865            let mut tt = IntraTxTypeCdfs::new();
866            let written = encode_coeffs(
867                &mut enc,
868                &mut cdfs,
869                size,
870                TxTypeCtx {
871                    set,
872                    intra_cdfs: &mut tt,
873                    intra_dir: 0,
874                    uv_mode,
875                    qindex_positive: true,
876                    lossless: false,
877                },
878                tx_type,
879                ptype,
880                az,
881                ds,
882                &levels,
883            )
884            .unwrap();
885            let data = enc.finish();
886            let mut dec = SymbolDecoder::new(&data, false).unwrap();
887            let mut cdfs = CoeffCdfs::new(2);
888            let mut tt = IntraTxTypeCdfs::new();
889            let read = decode_coeffs(
890                &mut dec,
891                &mut cdfs,
892                size,
893                TxTypeCtx {
894                    set,
895                    intra_cdfs: &mut tt,
896                    intra_dir: 0,
897                    uv_mode,
898                    qindex_positive: true,
899                    lossless: false,
900                },
901                ptype,
902                az,
903                ds,
904            )
905            .unwrap();
906            assert_eq!(read.eob, written.eob, "round {round}");
907            assert_eq!(read.quant[..w * h], written.quant[..w * h], "round {round}");
908            assert_eq!(read.quant[..w * h], levels[..], "round {round}: levels");
909            assert_eq!(
910                (read.cul_level, read.dc_category, read.tx_type),
911                (written.cul_level, written.dc_category, written.tx_type)
912            );
913        }
914    }
915
916    #[test]
917    fn base_eob_contexts_match_the_spec_buckets() {
918        let q = [0; MAX_COEFFS];
919        assert_eq!(
920            coeff_base_ctx(TxSize::Tx4x4, 0, &q, 0, 0, true),
921            SIG_COEF_CONTEXTS - 4
922        );
923        assert_eq!(
924            coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 1, true),
925            SIG_COEF_CONTEXTS - 3
926        );
927        assert_eq!(
928            coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 3, true),
929            SIG_COEF_CONTEXTS - 2
930        );
931        assert_eq!(
932            coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 9, true),
933            SIG_COEF_CONTEXTS - 1
934        );
935    }
936
937    #[test]
938    fn dc_position_base_context_is_zero() {
939        let q = [3; MAX_COEFFS];
940        assert_eq!(coeff_base_ctx(TxSize::Tx4x4, 0, &q, 0, 4, false), 0);
941    }
942
943    #[test]
944    fn base_context_folds_in_neighbour_magnitudes() {
945        // pos 5 -> row 1, col 1 (bwl 2). Put a 3 at pos 6 only: mag=3,
946        // ctx=(3+1)>>1=2, plus Coeff_Base_Ctx_Offset[TX_4X4][1][1]=6 -> 8.
947        let mut q = [0_i32; MAX_COEFFS];
948        q[6] = 3;
949        assert_eq!(coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 4, false), 8);
950    }
951
952    #[test]
953    fn br_context_at_dc_is_the_bare_magnitude() {
954        // pos 0: neighbours (0,1)=pos1, (1,0)=pos4, (1,1)=pos5. A single 5 at
955        // pos1 -> mag=min(5,15)=5, (5+1)>>1=3, pos==0 so ctx=3.
956        let mut q = [0_i32; MAX_COEFFS];
957        q[1] = 5;
958        assert_eq!(coeff_br_ctx(TxSize::Tx4x4, 0, &q, 0), 3);
959    }
960
961    #[test]
962    fn vertical_class_uses_position_offsets() {
963        // A 1D vertical transform (class 2) uses Coeff_Base_Pos_Ctx_Offset, not
964        // the 2D table: at pos 0 with no neighbours the ctx is the base offset.
965        let q = [0_i32; MAX_COEFFS];
966        // pos 8 in an 8x8 (bwl 3) -> row 1, col 0. No neighbours set, mag 0,
967        // ctx 0, idx=row=1 -> Coeff_Base_Pos_Ctx_Offset[1] = 31.
968        assert_eq!(
969            coeff_base_ctx(TxSize::Tx8x8, 2, &q, 8, 4, false),
970            (SIG_COEF_CONTEXTS_2D + 5) as usize
971        );
972    }
973
974    #[test]
975    fn scans_are_selected_by_size_and_type() {
976        assert_eq!(get_scan(TxSize::Tx4x4, TxType::DctDct).len(), 16);
977        assert_eq!(get_scan(TxSize::Tx8x8, TxType::DctDct).len(), 64);
978        // 64-wide sizes fall back to the 32x32 scan.
979        assert_eq!(get_scan(TxSize::Tx64x64, TxType::DctDct).len(), 1024);
980        assert_eq!(get_scan(TxSize::Tx16x64, TxType::DctDct).len(), 512);
981        // V_/H_ types switch to the row/column scans (compared by content:
982        // scan tables are `const`, so they have no stable address).
983        assert_eq!(get_scan(TxSize::Tx8x8, TxType::VDct), &MROW_SCAN_8X8[..]);
984        assert_eq!(get_scan(TxSize::Tx8x8, TxType::HDct), &MCOL_SCAN_8X8[..]);
985        // The default (2D) scan differs from the row scan for the same size.
986        assert_ne!(
987            get_scan(TxSize::Tx8x8, TxType::DctDct),
988            get_scan(TxSize::Tx8x8, TxType::VDct)
989        );
990    }
991
992    #[test]
993    fn an_all_zero_block_reads_one_symbol_and_stops() {
994        let data = [0x00; 8];
995        let mut dec = SymbolDecoder::new(&data, true).unwrap();
996        let mut cdfs = CoeffCdfs::new(0);
997        let mut intra = IntraTxTypeCdfs::new();
998        let tx = TxTypeCtx {
999            set: IntraTxSet::Set1,
1000            intra_cdfs: &mut intra,
1001            intra_dir: 0,
1002            uv_mode: 0,
1003            qindex_positive: false,
1004            lossless: true,
1005        };
1006        let block = decode_coeffs(&mut dec, &mut cdfs, TxSize::Tx4x4, tx, 0, 0, 0).unwrap();
1007        assert_eq!(block.tx_type, TxType::DctDct);
1008        if block.eob == 0 {
1009            assert_eq!(block.quant, [0; MAX_COEFFS]);
1010            assert_eq!(block.cul_level, 0);
1011        }
1012    }
1013
1014    #[test]
1015    fn golomb_reads_a_unary_prefix_then_data_bits() {
1016        let data = [0xFF; 4];
1017        let mut dec = SymbolDecoder::new(&data, true).unwrap();
1018        assert_eq!(read_golomb(&mut dec).unwrap(), 1);
1019    }
1020}