otf_pixels_codec_avif/av1/
symbol.rs1use super::bits::{BitReader, floor_log2};
22use otf_pixels_core::{PixelsError, Result};
23
24const EC_PROB_SHIFT: u32 = 6;
26const EC_MIN_PROB: u32 = 4;
28
29pub struct SymbolDecoder<'a> {
31 reader: BitReader<'a>,
32 value: u32,
35 range: u32,
37 max_bits: i64,
40 disable_cdf_update: bool,
42}
43
44impl<'a> SymbolDecoder<'a> {
45 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 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 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 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 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 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 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 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 #[must_use]
164 pub fn max_bits(&self) -> i64 {
165 self.max_bits
166 }
167}
168
169pub struct SymbolEncoder {
173 low: u64,
175 rng: u32,
177 cnt: i32,
179 precarry: Vec<u16>,
181 disable_cdf_update: bool,
182}
183
184impl SymbolEncoder {
185 #[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 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 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 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 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 #[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
292fn 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 *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 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 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 fn cdf_binary(c0: u16) -> [u16; 3] {
385 [c0, 1 << 15, 0]
386 }
387
388 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 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 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 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 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 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 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 assert!(cdf[0] > before, "{} !> {}", cdf[0], before);
478 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}