Skip to main content

otf_pixels_codec_avif/av1/
tx_size.rs

1//! Transform-size decode (spec §5.11.15, "read_tx_size", and §5.11.16,
2//! "read_block_tx_size").
3//!
4//! A coding block's luma transform size is read once per block for an intra
5//! frame. Lossless forces `TX_4X4`; otherwise the size starts at the largest
6//! rectangular transform that fits the block (`Max_Tx_Size_Rect`) and, when the
7//! frame is in `TX_MODE_SELECT`, a `tx_depth` symbol splits it down that many
8//! levels through `Split_Tx_Size`. The same size then applies to every transform
9//! block in the coding block (there is no variable-transform tree for intra).
10//!
11//! This module owns the pure size tables and the `tx_depth` symbol read with its
12//! neighbour-derived context. Wiring the resulting size through the reconstruct
13//! loop belongs with the tile decoder, which currently drives a `TX_4X4`-only
14//! path; the helpers here take the resolved neighbour widths directly.
15
16use super::cdf;
17use super::coder::{Site, TileCoder};
18use super::symbol::SymbolDecoder;
19use super::transform::TxSize;
20use otf_pixels_core::{PixelsError, Result};
21
22/// `BLOCK_4X4` (§6.10.4): the smallest block, index 0 of `BLOCK_SIZES`.
23pub const BLOCK_4X4: usize = 0;
24
25/// `Max_Tx_Size_Rect[BLOCK_SIZES]` (§9.3): the largest rectangular transform
26/// that fits each of the 22 block sizes.
27const MAX_TX_SIZE_RECT: [TxSize; 22] = [
28    TxSize::Tx4x4,
29    TxSize::Tx4x8,
30    TxSize::Tx8x4,
31    TxSize::Tx8x8,
32    TxSize::Tx8x16,
33    TxSize::Tx16x8,
34    TxSize::Tx16x16,
35    TxSize::Tx16x32,
36    TxSize::Tx32x16,
37    TxSize::Tx32x32,
38    TxSize::Tx32x64,
39    TxSize::Tx64x32,
40    TxSize::Tx64x64,
41    TxSize::Tx64x64,
42    TxSize::Tx64x64,
43    TxSize::Tx64x64,
44    TxSize::Tx4x16,
45    TxSize::Tx16x4,
46    TxSize::Tx8x32,
47    TxSize::Tx32x8,
48    TxSize::Tx16x64,
49    TxSize::Tx64x16,
50];
51
52/// `Max_Tx_Depth[BLOCK_SIZES]` (§9.3): how many times a block's transform must
53/// split to reach `TX_4X4`. Can exceed `MAX_TX_DEPTH`; the `tx_depth` symbol
54/// still only codes 0..=2, so deeper blocks cannot actually be coded that small.
55const MAX_TX_DEPTH: [usize; 22] = [
56    0, 1, 1, 1, 2, 2, 2, 3, 3, 3, 4, 4, 4, 4, 4, 4, 2, 2, 3, 3, 4, 4,
57];
58
59/// `Split_Tx_Size[TX_SIZES_ALL]` (§9.3): the transform size one depth down.
60const SPLIT_TX_SIZE: [TxSize; 19] = [
61    TxSize::Tx4x4,
62    TxSize::Tx4x4,
63    TxSize::Tx8x8,
64    TxSize::Tx16x16,
65    TxSize::Tx32x32,
66    TxSize::Tx4x4,
67    TxSize::Tx4x4,
68    TxSize::Tx8x8,
69    TxSize::Tx8x8,
70    TxSize::Tx16x16,
71    TxSize::Tx16x16,
72    TxSize::Tx32x32,
73    TxSize::Tx32x32,
74    TxSize::Tx4x8,
75    TxSize::Tx8x4,
76    TxSize::Tx8x16,
77    TxSize::Tx16x8,
78    TxSize::Tx16x32,
79    TxSize::Tx32x16,
80];
81
82/// `Max_Tx_Size_Rect[block]`.
83#[must_use]
84pub fn max_tx_size_rect(block: usize) -> TxSize {
85    MAX_TX_SIZE_RECT
86        .get(block)
87        .copied()
88        .unwrap_or(TxSize::Tx4x4)
89}
90
91/// `Max_Tx_Depth[block]`.
92#[must_use]
93pub fn max_tx_depth(block: usize) -> usize {
94    MAX_TX_DEPTH.get(block).copied().unwrap_or(0)
95}
96
97/// `Split_Tx_Size[txSz]`: one transform-depth step down.
98#[must_use]
99pub fn split_tx_size(tx: TxSize) -> TxSize {
100    SPLIT_TX_SIZE
101        .get(tx as usize)
102        .copied()
103        .unwrap_or(TxSize::Tx4x4)
104}
105
106/// The `BLOCK_SIZES` index for a block `w4 x h4` 4-sample units wide/high, or
107/// `None` if that is not a defined block shape.
108#[must_use]
109pub fn block_size_from_4x4(w4: usize, h4: usize) -> Option<usize> {
110    Some(match (w4, h4) {
111        (1, 1) => 0,
112        (1, 2) => 1,
113        (2, 1) => 2,
114        (2, 2) => 3,
115        (2, 4) => 4,
116        (4, 2) => 5,
117        (4, 4) => 6,
118        (4, 8) => 7,
119        (8, 4) => 8,
120        (8, 8) => 9,
121        (8, 16) => 10,
122        (16, 8) => 11,
123        (16, 16) => 12,
124        (16, 32) => 13,
125        (32, 16) => 14,
126        (32, 32) => 15,
127        (1, 4) => 16,
128        (4, 1) => 17,
129        (2, 8) => 18,
130        (8, 2) => 19,
131        (4, 16) => 20,
132        (16, 4) => 21,
133        _ => return None,
134    })
135}
136
137/// The `tx_depth` context (§8.3.2): whether the above/left neighbour transforms
138/// are at least as wide/tall as this block's maximum transform. `above_w` and
139/// `left_h` are the neighbour transform width/height in samples (0 when the
140/// neighbour is unavailable), as `get_above_tx_width` / `get_left_tx_height`
141/// resolve them for an intra block.
142#[must_use]
143pub fn tx_depth_ctx(above_w: usize, left_h: usize, max_rect: TxSize) -> usize {
144    usize::from(above_w >= max_rect.width()) + usize::from(left_h >= max_rect.height())
145}
146
147/// The adapting `tx_depth` CDFs for one tile, one per maximum-transform category.
148pub struct TxDepthCdfs {
149    tx8x8: [[u16; 3]; 3],
150    tx16x16: [[u16; 4]; 3],
151    tx32x32: [[u16; 4]; 3],
152    tx64x64: [[u16; 4]; 3],
153}
154
155impl TxDepthCdfs {
156    /// Clone the frame defaults (these CDFs do not depend on the quantiser).
157    #[must_use]
158    pub fn new() -> Self {
159        Self {
160            tx8x8: cdf::DEFAULT_TX_8X8_CDF,
161            tx16x16: cdf::DEFAULT_TX_16X16_CDF,
162            tx32x32: cdf::DEFAULT_TX_32X32_CDF,
163            tx64x64: cdf::DEFAULT_TX_64X64_CDF,
164        }
165    }
166}
167
168impl Default for TxDepthCdfs {
169    fn default() -> Self {
170        Self::new()
171    }
172}
173
174fn row_mut<T>(slice: &mut [T], index: usize) -> Result<&mut T> {
175    slice
176        .get_mut(index)
177        .ok_or_else(|| PixelsError::malformed("avif", "an AV1 tx-size CDF index ran out of range"))
178}
179
180/// The inputs to [`read_tx_size`] other than the decoder and CDFs.
181pub struct TxSizeParams {
182    /// The coding block's `BLOCK_SIZES` index.
183    pub block: usize,
184    /// `TxMode == TX_MODE_SELECT`: the frame codes a transform depth per block.
185    pub tx_mode_select: bool,
186    /// The lossless flag; forces `TX_4X4` and reads no symbol.
187    pub lossless: bool,
188    /// The caller's selection gate (`!skip || !is_inter`).
189    pub allow_select: bool,
190    /// The above neighbour's transform width in samples (0 if unavailable).
191    pub above_w: usize,
192    /// The left neighbour's transform height in samples (0 if unavailable).
193    pub left_h: usize,
194}
195
196/// Resolve a coding block's luma transform size (`read_tx_size`, §5.11.15).
197///
198/// A `tx_depth` symbol is read only when the block is larger than 4x4 and
199/// selection is active; otherwise the size is `Max_Tx_Size_Rect` (or `TX_4X4`
200/// for lossless). See [`TxSizeParams`] for the inputs.
201///
202/// # Errors
203///
204/// Propagates arithmetic-decoder errors and rejects an out-of-range context.
205pub fn read_tx_size(
206    dec: &mut SymbolDecoder<'_>,
207    cdfs: &mut TxDepthCdfs,
208    params: &TxSizeParams,
209) -> Result<TxSize> {
210    code_tx_size(dec, cdfs, params)
211}
212
213/// [`read_tx_size`] over any [`TileCoder`], for the tile syntax to share
214/// between decoding and encoding.
215pub(crate) fn code_tx_size(
216    dec: &mut impl TileCoder,
217    cdfs: &mut TxDepthCdfs,
218    params: &TxSizeParams,
219) -> Result<TxSize> {
220    if params.lossless {
221        return Ok(TxSize::Tx4x4);
222    }
223    let max_rect = max_tx_size_rect(params.block);
224    let max_depth = max_tx_depth(params.block);
225    let mut tx = max_rect;
226    if params.block > BLOCK_4X4 && params.allow_select && params.tx_mode_select {
227        let ctx = tx_depth_ctx(params.above_w, params.left_h, max_rect);
228        let depth = match max_depth {
229            4 => dec.symbol(row_mut(&mut cdfs.tx64x64, ctx)?, Site::TxDepth)?,
230            3 => dec.symbol(row_mut(&mut cdfs.tx32x32, ctx)?, Site::TxDepth)?,
231            2 => dec.symbol(row_mut(&mut cdfs.tx16x16, ctx)?, Site::TxDepth)?,
232            _ => dec.symbol(row_mut(&mut cdfs.tx8x8, ctx)?, Site::TxDepth)?,
233        };
234        for _ in 0..depth {
235            tx = split_tx_size(tx);
236        }
237    }
238    Ok(tx)
239}
240
241#[cfg(test)]
242#[allow(
243    clippy::unwrap_used,
244    clippy::indexing_slicing,
245    clippy::panic,
246    reason = "tests operate on known-good values and assert shapes directly"
247)]
248mod tests {
249    use super::*;
250
251    #[test]
252    fn tables_agree_with_the_spec() {
253        // 16x16 (block 6) takes a 16x16 transform, splitting to 8x8 then 4x4.
254        assert_eq!(max_tx_size_rect(6), TxSize::Tx16x16);
255        assert_eq!(max_tx_depth(6), 2);
256        assert_eq!(split_tx_size(TxSize::Tx16x16), TxSize::Tx8x8);
257        assert_eq!(split_tx_size(TxSize::Tx8x8), TxSize::Tx4x4);
258        assert_eq!(split_tx_size(TxSize::Tx4x4), TxSize::Tx4x4);
259        // 64x64 caps the depth at 4 and splits square down one level.
260        assert_eq!(max_tx_size_rect(12), TxSize::Tx64x64);
261        assert_eq!(max_tx_depth(12), 4);
262        assert_eq!(split_tx_size(TxSize::Tx64x64), TxSize::Tx32x32);
263        // A rectangle splits along its long side first.
264        assert_eq!(split_tx_size(TxSize::Tx4x16), TxSize::Tx4x8);
265    }
266
267    #[test]
268    fn block_size_lookup_round_trips_the_defined_shapes() {
269        assert_eq!(block_size_from_4x4(1, 1), Some(0));
270        assert_eq!(block_size_from_4x4(4, 4), Some(6));
271        assert_eq!(block_size_from_4x4(16, 4), Some(21));
272        // 4x32 (w4=1, h4=8) is not a defined AV1 block shape.
273        assert_eq!(block_size_from_4x4(1, 8), None);
274    }
275
276    fn params(block: usize, tx_mode_select: bool, lossless: bool) -> TxSizeParams {
277        TxSizeParams {
278            block,
279            tx_mode_select,
280            lossless,
281            allow_select: true,
282            above_w: 0,
283            left_h: 0,
284        }
285    }
286
287    #[test]
288    fn lossless_is_always_4x4_and_reads_nothing() {
289        let data = [0xFF; 4];
290        let mut dec = SymbolDecoder::new(&data, true).unwrap();
291        let mut cdfs = TxDepthCdfs::new();
292        // Lossless: TX_4X4 regardless of block or mode, no symbol consumed.
293        let tx = read_tx_size(&mut dec, &mut cdfs, &params(12, true, true)).unwrap();
294        assert_eq!(tx, TxSize::Tx4x4);
295    }
296
297    #[test]
298    fn without_selection_the_size_is_the_max_rect() {
299        let data = [0xFF; 4];
300        let mut dec = SymbolDecoder::new(&data, true).unwrap();
301        let mut cdfs = TxDepthCdfs::new();
302        // TX_MODE_SELECT off: the largest transform, no symbol read.
303        let tx = read_tx_size(&mut dec, &mut cdfs, &params(6, false, false)).unwrap();
304        assert_eq!(tx, TxSize::Tx16x16);
305        // A 4x4 block is always TX_4X4 even with selection on (block == BLOCK_4X4).
306        let tx = read_tx_size(&mut dec, &mut cdfs, &params(BLOCK_4X4, true, false)).unwrap();
307        assert_eq!(tx, TxSize::Tx4x4);
308    }
309
310    #[test]
311    fn a_selected_size_is_the_max_rect_or_a_split_of_it() {
312        let data = [0x80, 0x00, 0x00, 0x00, 0x00, 0x00];
313        let mut dec = SymbolDecoder::new(&data, true).unwrap();
314        let mut cdfs = TxDepthCdfs::new();
315        // 16x16 with selection: the result is 16x16, 8x8, or 4x4 depending on the
316        // depth symbol — every reachable size on the split chain.
317        let tx = read_tx_size(&mut dec, &mut cdfs, &params(6, true, false)).unwrap();
318        assert!(matches!(
319            tx,
320            TxSize::Tx16x16 | TxSize::Tx8x8 | TxSize::Tx4x4
321        ));
322    }
323
324    #[test]
325    fn depth_context_counts_the_larger_neighbours() {
326        // 16x16 max transform is 16 wide / 16 tall.
327        let m = TxSize::Tx16x16;
328        assert_eq!(tx_depth_ctx(0, 0, m), 0); // both neighbours smaller/absent
329        assert_eq!(tx_depth_ctx(16, 0, m), 1); // above at least as wide
330        assert_eq!(tx_depth_ctx(32, 16, m), 2); // both at least as large
331    }
332}