Skip to main content

otf_pixels_codec_avif/av1/
symbol.rs

1//! The AV1 multi-symbol arithmetic decoder (spec §8.2.3–8.2.6).
2//!
3//! This is the entropy engine every tile symbol flows through. It is a range
4//! decoder over 15-bit cumulative distribution functions, with the twist that
5//! AV1 works with the *inverse* CDF (`f = 32768 - cdf[i]`) and gives every
6//! symbol a floor probability of [`EC_MIN_PROB`] so the range never collapses.
7//! Each `read_symbol` optionally adapts the CDF toward the symbol it just
8//! decoded, which is why the CDF is passed by mutable reference.
9//!
10//! The transcription is taken verbatim from the AV1 specification's symbol
11//! decoding process and cross-checked against libaom's `entdec.c`; the two
12//! agree on the inverse-CDF arithmetic and on the renormalization
13//! (`SymbolValue = paddedData ^ (((SymbolValue + 1) << bits) - 1)`). Numeric
14//! conformance against real streams is established by the libaom differential
15//! harness, which is the only trustworthy check for an entropy coder.
16//!
17//! Every read is fallible: a stream that runs out mid-renormalization is a
18//! returned error, and `SymbolMaxBits` accounting means the padding-zero region
19//! past the real bytes is entered deliberately, never by reading off the end.
20
21use super::bits::{BitReader, floor_log2};
22use otf_pixels_core::{PixelsError, Result};
23
24/// Bits of CDF precision dropped during the range update (§3, `EC_PROB_SHIFT`).
25const EC_PROB_SHIFT: u32 = 6;
26/// The floor probability every symbol is guaranteed (§3, `EC_MIN_PROB`).
27const EC_MIN_PROB: u32 = 4;
28
29/// A range decoder over an AV1 tile's symbol data.
30pub struct SymbolDecoder<'a> {
31    reader: BitReader<'a>,
32    /// `SymbolValue` — the decoded position within the current range, stored in
33    /// the spec's inverted form.
34    value: u32,
35    /// `SymbolRange` — the current range, always in `[2^15, 2^16)`.
36    range: u32,
37    /// `SymbolMaxBits` — real bits still available. Goes negative once the
38    /// decoder enters the implicit zero-padding past the end of the data.
39    max_bits: i64,
40    /// Whether per-symbol CDF adaptation is switched off for the frame.
41    disable_cdf_update: bool,
42}
43
44impl<'a> SymbolDecoder<'a> {
45    /// Initialise the decoder over a tile's `sz`-byte symbol partition
46    /// (`init_symbol`, §8.2.2). The position must be byte-aligned, which the
47    /// syntax guarantees at every call site.
48    pub fn new(data: &'a [u8], disable_cdf_update: bool) -> Result<Self> {
49        let sz = data.len();
50        let mut reader = BitReader::new(data);
51        let num_bits = u32::try_from(sz.saturating_mul(8).min(15)).unwrap_or(15);
52        let buf = reader.f(num_bits)?;
53        let padded_buf = buf << (15 - num_bits);
54        let value = ((1_u32 << 15) - 1) ^ padded_buf;
55        let max_bits = (8 * sz as i64) - 15;
56        Ok(Self {
57            reader,
58            value,
59            range: 1 << 15,
60            max_bits,
61            disable_cdf_update,
62        })
63    }
64
65    /// Decode one symbol against `cdf`, adapting it unless updates are disabled
66    /// (`read_symbol`, §8.2.6). `cdf` has `N + 1` entries: `N` cumulative
67    /// frequencies with `cdf[N-1] == 1 << 15`, then an adaptation counter.
68    pub fn read_symbol(&mut self, cdf: &mut [u16]) -> Result<usize> {
69        let len = cdf.len();
70        let Some(n) = len.checked_sub(1).filter(|&n| n >= 1) else {
71            return Err(PixelsError::malformed(
72                "avif",
73                "an AV1 CDF must hold at least one symbol and a counter",
74            ));
75        };
76
77        // decode_symbol: walk the inverse CDF until SymbolValue lands in a
78        // symbol's interval. cur decreases as the symbol index rises, so the
79        // first index whose cur is at or below the value is the answer.
80        let mut cur = self.range;
81        let mut prev = cur;
82        let mut symbol = n - 1;
83        for (k, &c) in cdf.iter().take(n).enumerate() {
84            prev = cur;
85            let f = (1_u32 << 15) - u32::from(c);
86            cur = ((self.range >> 8) * (f >> EC_PROB_SHIFT)) >> (7 - EC_PROB_SHIFT);
87            cur += EC_MIN_PROB * (n as u32 - 1 - k as u32);
88            if self.value >= cur {
89                symbol = k;
90                break;
91            }
92        }
93
94        self.range = prev - cur;
95        self.value -= cur;
96        self.renormalize()?;
97
98        if !self.disable_cdf_update {
99            update_cdf(cdf, symbol, n);
100        }
101        Ok(symbol)
102    }
103
104    /// Renormalize the range and refill the value with new bits (§8.2.6 steps
105    /// 1–7). `bits` new bits enter; past the real data they are implicit zeros.
106    fn renormalize(&mut self) -> Result<()> {
107        let bits = 15 - floor_log2(self.range);
108        self.range <<= bits;
109        let available = self.max_bits.max(0);
110        let num_bits = if i64::from(bits) < available {
111            bits
112        } else {
113            // available < bits <= 15, so it fits in u32.
114            available as u32
115        };
116        let new_data = self.reader.f(num_bits)?;
117        let padded_data = new_data << (bits - num_bits);
118        self.value = padded_data ^ (((self.value + 1) << bits) - 1);
119        self.max_bits -= i64::from(bits);
120        Ok(())
121    }
122
123    /// Decode one equiprobable bit (`read_bool`, §8.2.3). The transient CDF's
124    /// adaptation is never observed, so it is skipped.
125    pub fn read_bool(&mut self) -> Result<bool> {
126        let mut cdf = [1_u16 << 14, 1_u16 << 15, 0];
127        let saved = self.disable_cdf_update;
128        self.disable_cdf_update = true;
129        let symbol = self.read_symbol(&mut cdf);
130        self.disable_cdf_update = saved;
131        Ok(symbol? != 0)
132    }
133
134    /// Decode an `n`-bit literal, most-significant bit first (`read_literal`,
135    /// §8.2.5).
136    pub fn read_literal(&mut self, n: u32) -> Result<u32> {
137        let mut x = 0;
138        for _ in 0..n {
139            x = 2 * x + u32::from(self.read_bool()?);
140        }
141        Ok(x)
142    }
143
144    /// Decode a non-symmetric `NS(n)` value in `0..n` (`read_ns`, §8.2.4). Used
145    /// for the palette colour-index map's first sample, which is uniform over
146    /// the palette size.
147    pub fn read_ns(&mut self, n: u32) -> Result<u32> {
148        if n <= 1 {
149            return Ok(0);
150        }
151        let w = floor_log2(n) + 1;
152        let m = (1 << w) - n;
153        let v = self.read_literal(w - 1)?;
154        if v < m {
155            return Ok(v);
156        }
157        let extra = self.read_literal(1)?;
158        Ok((v << 1) - m + extra)
159    }
160
161    /// The number of real bits still available (may be negative once the
162    /// decoder is in the padding region). Exposed for the tile-exit checks.
163    #[must_use]
164    pub fn max_bits(&self) -> i64 {
165        self.max_bits
166    }
167}
168
169/// The multi-symbol arithmetic encoder: libaom's `od_ec_enc`, the exact
170/// inverse of [`SymbolDecoder`]. Symbols are coded against the same CDFs,
171/// adapted the same way, so a stream it writes decodes symbol for symbol.
172pub struct SymbolEncoder {
173    /// The low end of the current interval, below the bits already flushed.
174    low: u64,
175    /// The interval's size, in `[2^15, 2^16)` between symbols.
176    rng: u32,
177    /// Bits buffered in `low` beyond a whole byte, offset as libaom keeps it.
178    cnt: i32,
179    /// Output bytes before carry propagation, each with room for a carry.
180    precarry: Vec<u16>,
181    disable_cdf_update: bool,
182}
183
184impl SymbolEncoder {
185    /// A fresh encoder for one tile.
186    #[must_use]
187    pub const fn new(disable_cdf_update: bool) -> Self {
188        Self {
189            low: 0,
190            rng: 0x8000,
191            cnt: -9,
192            precarry: Vec::new(),
193            disable_cdf_update,
194        }
195    }
196
197    /// Encode `symbol` against `cdf`, adapting it as the decoder will.
198    pub fn write_symbol(&mut self, cdf: &mut [u16], symbol: usize) {
199        let n = cdf.len().saturating_sub(1).max(1);
200        let symbol = symbol.min(n - 1);
201        let r = self.rng;
202        let bound = |k: usize| -> u32 {
203            let f = (1_u32 << 15) - u32::from(cdf.get(k).copied().unwrap_or(1 << 15));
204            (((r >> 8) * (f >> EC_PROB_SHIFT)) >> (7 - EC_PROB_SHIFT))
205                + EC_MIN_PROB * (n as u32 - 1 - k as u32)
206        };
207        let v = bound(symbol);
208        let (low, rng) = if symbol > 0 {
209            let u = bound(symbol - 1);
210            (self.low + u64::from(r - u), u - v)
211        } else {
212            (self.low, r - v)
213        };
214        self.normalize(low, rng);
215        if !self.disable_cdf_update {
216            update_cdf(cdf, symbol, n);
217        }
218    }
219
220    /// One equiprobable bit (`read_bool`'s inverse).
221    pub fn write_bool(&mut self, bit: bool) {
222        let mut cdf = [1_u16 << 14, 1_u16 << 15, 0];
223        let saved = self.disable_cdf_update;
224        self.disable_cdf_update = true;
225        self.write_symbol(&mut cdf, usize::from(bit));
226        self.disable_cdf_update = saved;
227    }
228
229    /// An `n`-bit literal, most-significant bit first.
230    pub fn write_literal(&mut self, n: u32, value: u32) {
231        for i in (0..n).rev() {
232            self.write_bool((value >> i) & 1 == 1);
233        }
234    }
235
236    /// `od_ec_enc_normalize`: shift the interval back up to 16 bits, moving
237    /// whole bytes of `low` out to the pre-carry buffer.
238    fn normalize(&mut self, mut low: u64, rng: u32) {
239        let d = 15 - floor_log2(rng) as i32;
240        let mut c = self.cnt;
241        let mut s = c + d;
242        if s >= 0 {
243            c += 16;
244            let mut m = (1_u64 << c) - 1;
245            if s >= 8 {
246                self.precarry.push((low >> c) as u16);
247                low &= m;
248                c -= 8;
249                m >>= 8;
250            }
251            self.precarry.push((low >> c) as u16);
252            s = c + d - 24;
253            low &= m;
254        }
255        self.low = low << d;
256        self.rng = rng << d;
257        self.cnt = s;
258    }
259
260    /// `od_ec_enc_done`: flush the fewest bits that decode correctly whatever
261    /// follows (ending in the trailing one bit), then propagate carries.
262    #[must_use]
263    pub fn finish(mut self) -> Vec<u8> {
264        let mut c = self.cnt;
265        let mut s = 10 + c;
266        let m = 0x3fff_u64;
267        let mut e = ((self.low + m) & !m) | (m + 1);
268        if s > 0 {
269            let mut n = (1_u64 << (c + 16)) - 1;
270            loop {
271                self.precarry.push((e >> (c + 16)) as u16);
272                e &= n;
273                s -= 8;
274                c -= 8;
275                n >>= 8;
276                if s <= 0 {
277                    break;
278                }
279            }
280        }
281        let mut out = vec![0_u8; self.precarry.len()];
282        let mut carry = 0_u32;
283        for (slot, &v) in out.iter_mut().zip(&self.precarry).rev() {
284            carry += u32::from(v);
285            *slot = carry as u8;
286            carry >>= 8;
287        }
288        out
289    }
290}
291
292/// Adapt `cdf` toward `symbol` (`update_cdf`, §8.2.6). The counter at `cdf[n]`
293/// slows adaptation as a symbol is seen more often.
294fn update_cdf(cdf: &mut [u16], symbol: usize, n: usize) {
295    let count = cdf.get(n).copied().unwrap_or(0);
296    let rate = 3 + u32::from(count > 15) + u32::from(count > 31) + floor_log2(n as u32).min(2);
297    let mut tmp: u32 = 0;
298    for (i, slot) in cdf.iter_mut().take(n.saturating_sub(1)).enumerate() {
299        if i == symbol {
300            tmp = 1 << 15;
301        }
302        let ci = u32::from(*slot);
303        let updated = if tmp < ci {
304            ci - ((ci - tmp) >> rate)
305        } else {
306            ci + ((tmp - ci) >> rate)
307        };
308        // updated stays within [0, 1<<15], so the cast never truncates.
309        *slot = updated as u16;
310    }
311    if let Some(counter) = cdf.get_mut(n) {
312        if *counter < 32 {
313            *counter += 1;
314        }
315    }
316}
317
318#[cfg(test)]
319#[allow(
320    clippy::unwrap_used,
321    clippy::indexing_slicing,
322    clippy::panic,
323    reason = "tests operate on known-good values and assert shapes directly"
324)]
325mod tests {
326    use super::*;
327
328    #[test]
329    fn the_encoder_round_trips_through_the_decoder() {
330        // Pseudo-random symbols over CDFs of several sizes, adapted on both
331        // sides, interleaved with literals: every symbol must come back.
332        let mut state = 0x2545_f491_u32;
333        let mut next = || {
334            state ^= state << 13;
335            state ^= state >> 17;
336            state ^= state << 5;
337            state
338        };
339        for &adapt in &[true, false] {
340            let make = |n: usize| -> Vec<u16> {
341                let mut cdf: Vec<u16> = (1..=n).map(|k| ((k * 32768) / n) as u16).collect();
342                cdf.push(0);
343                cdf
344            };
345            let sizes = [2_usize, 3, 4, 8, 13, 16];
346            let mut enc_cdfs: Vec<Vec<u16>> = sizes.iter().map(|&n| make(n)).collect();
347            let mut dec_cdfs = enc_cdfs.clone();
348            let mut plan = Vec::new();
349            let mut enc = SymbolEncoder::new(!adapt);
350            for _ in 0..5000 {
351                let which = (next() % sizes.len() as u32) as usize;
352                // Skew towards low symbols so adaptation actually moves.
353                let r = next();
354                let symbol = if r % 4 == 0 {
355                    (r as usize >> 8) % sizes[which]
356                } else {
357                    0
358                };
359                if next() % 7 == 0 {
360                    let v = next() % 64;
361                    enc.write_literal(6, v);
362                    plan.push((usize::MAX, v as usize));
363                } else {
364                    enc.write_symbol(&mut enc_cdfs[which], symbol);
365                    plan.push((which, symbol));
366                }
367            }
368            let data = enc.finish();
369            let mut dec = SymbolDecoder::new(&data, !adapt).unwrap();
370            for (i, &(which, value)) in plan.iter().enumerate() {
371                let got = if which == usize::MAX {
372                    dec.read_literal(6).unwrap() as usize
373                } else {
374                    dec.read_symbol(&mut dec_cdfs[which]).unwrap()
375                };
376                assert_eq!(got, value, "symbol {i} (adapt {adapt})");
377            }
378            assert_eq!(enc_cdfs, dec_cdfs, "adaptation diverged");
379        }
380    }
381
382    /// A binary CDF: `cdf[0]` is P(symbol 0), then the mandatory `1<<15` and the
383    /// adaptation counter.
384    fn cdf_binary(c0: u16) -> [u16; 3] {
385        [c0, 1 << 15, 0]
386    }
387
388    /// A genuine three-symbol CDF (two cumulative splits + counter).
389    fn cdf3(c0: u16, c1: u16) -> [u16; 4] {
390        [c0, c1, 1 << 15, 0]
391    }
392
393    #[test]
394    fn init_state_matches_the_spec() {
395        let data = [0xAB, 0xCD, 0xEF];
396        let dec = SymbolDecoder::new(&data, false).unwrap();
397        assert_eq!(dec.range, 1 << 15);
398        // numBits = min(24,15) = 15; buf = top 15 bits of 0xABCD. paddedBuf is
399        // buf << (15 - 15) = buf.
400        let padded_buf = u32::from(0xABCD_u16) >> 1;
401        assert_eq!(dec.value, ((1 << 15) - 1) ^ padded_buf);
402        assert_eq!(dec.max_bits, 8 * 3 - 15);
403    }
404
405    #[test]
406    fn a_tiny_buffer_starts_in_the_padding_region() {
407        // One byte: only 8 real bits, so SymbolMaxBits is negative from the
408        // start and renormalization reads no further real bits.
409        let dec = SymbolDecoder::new(&[0x00], false).unwrap();
410        assert_eq!(dec.max_bits, 8 - 15);
411    }
412
413    #[test]
414    fn a_cdf_certain_of_the_first_symbol_decodes_it_when_the_value_is_high() {
415        // An all-zero buffer makes init's SymbolValue = 0x7FFF ^ 0 = 0x7FFF,
416        // which sits in the high interval. With cdf[0] = 32767 the low-value
417        // sliver belongs to symbol 1, so a high value pins the result to 0, and
418        // the all-zero stream keeps it there.
419        let mut dec = SymbolDecoder::new(&[0x00; 6], true).unwrap();
420        for _ in 0..8 {
421            assert_eq!(dec.read_symbol(&mut cdf_binary(32767)).unwrap(), 0);
422        }
423    }
424
425    #[test]
426    fn a_cdf_certain_of_the_last_symbol_decodes_it_when_the_value_is_low() {
427        // An all-0xFF buffer makes init's SymbolValue = 0x7FFF ^ 0x7FFF = 0,
428        // the lowest value. With cdf[0] = 1 almost the entire range belongs to
429        // symbol 1, and value 0 lands squarely in it.
430        let mut dec = SymbolDecoder::new(&[0xFF; 6], true).unwrap();
431        for _ in 0..8 {
432            assert_eq!(dec.read_symbol(&mut cdf_binary(1)).unwrap(), 1);
433        }
434    }
435
436    #[test]
437    fn read_literal_composes_read_bool() {
438        // read_literal(n) must equal n read_bool calls, MSB first, on the same
439        // stream. Run each on its own decoder over identical data.
440        let data = [0x3C, 0xA7, 0x91, 0x08, 0x55];
441        let mut a = SymbolDecoder::new(&data, false).unwrap();
442        let literal = a.read_literal(5).unwrap();
443
444        let mut b = SymbolDecoder::new(&data, false).unwrap();
445        let mut composed = 0;
446        for _ in 0..5 {
447            composed = 2 * composed + u32::from(b.read_bool().unwrap());
448        }
449        assert_eq!(literal, composed);
450    }
451
452    #[test]
453    fn decoding_is_deterministic_for_the_same_input() {
454        let data = [0x9E, 0x42, 0x17, 0xCB, 0x30, 0x8A];
455        let decode_all = || {
456            let mut dec = SymbolDecoder::new(&data, false).unwrap();
457            let mut out = Vec::new();
458            for _ in 0..12 {
459                out.push(dec.read_symbol(&mut cdf3(1 << 13, 3 << 13)).unwrap());
460            }
461            out
462        };
463        assert_eq!(decode_all(), decode_all());
464    }
465
466    #[test]
467    fn adaptation_moves_the_cdf_toward_the_decoded_symbol() {
468        // A high initial value (zero buffer) decodes symbol 0 here; confirm
469        // cdf[0] climbs toward 1<<15 as that symbol is reinforced.
470        let data = [0x00, 0x00, 0x00, 0x00, 0x00, 0x00];
471        let mut dec = SymbolDecoder::new(&data, false).unwrap();
472        let mut cdf = cdf_binary(1 << 14);
473        let before = cdf[0];
474        let symbol = dec.read_symbol(&mut cdf).unwrap();
475        assert_eq!(symbol, 0);
476        // Seeing symbol 0 pushes cdf[0] upward (toward certainty of 0).
477        assert!(cdf[0] > before, "{} !> {}", cdf[0], before);
478        // The counter advanced.
479        assert_eq!(cdf[2], 1);
480    }
481
482    #[test]
483    fn a_malformed_cdf_is_rejected_not_panicked() {
484        let mut dec = SymbolDecoder::new(&[0x00, 0x11], false).unwrap();
485        let mut too_short = [1_u16 << 15];
486        assert!(dec.read_symbol(&mut too_short).is_err());
487    }
488}