use super::cdf;
use super::symbol::{SymbolDecoder, SymbolEncoder};
use super::transform::{TxSize, TxType};
use super::transform_type::{
IntraTxSet, IntraTxTypeCdfs, chroma_tx_type, read_transform_type, write_transform_type,
};
use otf_pixels_core::{PixelsError, Result};
include!("scan_tables.rs");
fn pick4<T: Copy>(arr: [T; 4], q: usize) -> T {
let [a, b, c, d] = arr;
match q {
1 => b,
2 => c,
3 => d,
_ => a,
}
}
fn cdf_row<T>(slice: &mut [T], index: usize) -> Result<&mut T> {
slice.get_mut(index).ok_or_else(|| {
PixelsError::malformed("avif", "an AV1 coefficient CDF index ran out of range")
})
}
const NUM_BASE_LEVELS: i32 = 2;
const COEFF_BASE_RANGE: i32 = 12;
const BR_CDF_SIZE: i32 = 4;
const SIG_COEF_CONTEXTS: usize = 42;
const SIG_COEF_CONTEXTS_2D: i32 = 26;
const SIG_COEF_CONTEXTS_EOB: usize = 4;
const MAX_COEFFS: usize = 1024;
fn tx_class(tx_type: TxType) -> usize {
match tx_type {
TxType::VDct | TxType::VAdst | TxType::VFlipadst => 2,
TxType::HDct | TxType::HAdst | TxType::HFlipadst => 1,
_ => 0,
}
}
const SIG_REF_DIFF_OFFSET: [[(i32, i32); 5]; 3] = [
[(0, 1), (1, 0), (1, 1), (0, 2), (2, 0)],
[(0, 1), (1, 0), (0, 2), (0, 3), (0, 4)],
[(0, 1), (1, 0), (2, 0), (3, 0), (4, 0)],
];
const MAG_REF_OFFSET: [[(i32, i32); 3]; 3] = [
[(0, 1), (1, 0), (1, 1)],
[(0, 1), (1, 0), (0, 2)],
[(0, 1), (1, 0), (2, 0)],
];
const COEFF_BASE_POS_CTX_OFFSET: [i32; 3] = [
SIG_COEF_CONTEXTS_2D,
SIG_COEF_CONTEXTS_2D + 5,
SIG_COEF_CONTEXTS_2D + 10,
];
const CBO_4X4: [[i32; 5]; 5] = [
[0, 1, 6, 6, 0],
[1, 6, 6, 21, 0],
[6, 6, 21, 21, 0],
[6, 21, 21, 21, 0],
[0, 0, 0, 0, 0],
];
const CBO_SQR: [[i32; 5]; 5] = [
[0, 1, 6, 6, 21],
[1, 6, 6, 21, 21],
[6, 6, 21, 21, 21],
[6, 21, 21, 21, 21],
[21, 21, 21, 21, 21],
];
const CBO_NARROW: [[i32; 5]; 5] = [
[0, 11, 11, 11, 0],
[11, 11, 11, 11, 0],
[6, 6, 21, 21, 0],
[6, 21, 21, 21, 0],
[21, 21, 21, 21, 0],
];
const CBO_SHORT: [[i32; 5]; 5] = [
[0, 16, 6, 6, 21],
[16, 16, 6, 21, 21],
[16, 16, 21, 21, 21],
[16, 16, 21, 21, 21],
[0, 0, 0, 0, 0],
];
const CBO_TALL: [[i32; 5]; 5] = [
[0, 11, 11, 11, 11],
[11, 11, 11, 11, 11],
[6, 6, 21, 21, 21],
[6, 21, 21, 21, 21],
[21, 21, 21, 21, 21],
];
const CBO_WIDE: [[i32; 5]; 5] = [
[0, 16, 6, 6, 21],
[16, 16, 6, 21, 21],
[16, 16, 21, 21, 21],
[16, 16, 21, 21, 21],
[16, 16, 21, 21, 21],
];
const COEFF_BASE_CTX_OFFSET: [[[i32; 5]; 5]; 19] = [
CBO_4X4, CBO_SQR, CBO_SQR, CBO_SQR, CBO_SQR, CBO_NARROW, CBO_SHORT, CBO_TALL, CBO_WIDE,
CBO_TALL, CBO_WIDE, CBO_TALL, CBO_WIDE, CBO_NARROW, CBO_SHORT, CBO_TALL, CBO_WIDE, CBO_TALL,
CBO_WIDE,
];
pub struct CoeffCdfs {
txb_skip: [[[u16; 3]; 13]; 5],
eob_pt_16: [[[u16; 6]; 2]; 2],
eob_pt_32: [[[u16; 7]; 2]; 2],
eob_pt_64: [[[u16; 8]; 2]; 2],
eob_pt_128: [[[u16; 9]; 2]; 2],
eob_pt_256: [[[u16; 10]; 2]; 2],
eob_pt_512: [[u16; 11]; 2],
eob_pt_1024: [[u16; 12]; 2],
eob_extra: [[[[u16; 3]; 9]; 2]; 5],
coeff_base_eob: [[[[u16; 4]; 4]; 2]; 5],
coeff_base: [[[[u16; 5]; 42]; 2]; 5],
coeff_br: [[[[u16; 5]; 21]; 2]; 5],
dc_sign: [[[u16; 3]; 3]; 2],
}
impl CoeffCdfs {
#[must_use]
pub fn new(qctx: usize) -> Self {
let q = qctx.min(3);
Self {
txb_skip: pick4(cdf::DEFAULT_TXB_SKIP_CDF, q),
eob_pt_16: pick4(cdf::DEFAULT_EOB_PT_16_CDF, q),
eob_pt_32: pick4(cdf::DEFAULT_EOB_PT_32_CDF, q),
eob_pt_64: pick4(cdf::DEFAULT_EOB_PT_64_CDF, q),
eob_pt_128: pick4(cdf::DEFAULT_EOB_PT_128_CDF, q),
eob_pt_256: pick4(cdf::DEFAULT_EOB_PT_256_CDF, q),
eob_pt_512: pick4(cdf::DEFAULT_EOB_PT_512_CDF, q),
eob_pt_1024: pick4(cdf::DEFAULT_EOB_PT_1024_CDF, q),
eob_extra: pick4(cdf::DEFAULT_EOB_EXTRA_CDF, q),
coeff_base_eob: pick4(cdf::DEFAULT_COEFF_BASE_EOB_CDF, q),
coeff_base: pick4(cdf::DEFAULT_COEFF_BASE_CDF, q),
coeff_br: pick4(cdf::DEFAULT_COEFF_BR_CDF, q),
dc_sign: pick4(cdf::DEFAULT_DC_SIGN_CDF, q),
}
}
}
pub struct CoeffBlock {
pub quant: [i32; MAX_COEFFS],
pub eob: usize,
pub cul_level: u8,
pub dc_category: u8,
pub tx_type: TxType,
}
pub struct TxTypeCtx<'a> {
pub set: IntraTxSet,
pub intra_cdfs: &'a mut IntraTxTypeCdfs,
pub intra_dir: usize,
pub uv_mode: usize,
pub qindex_positive: bool,
pub lossless: bool,
}
fn get_scan(tx_size: TxSize, tx_type: TxType) -> &'static [u16] {
match tx_size {
TxSize::Tx16x64 => return &DEFAULT_SCAN_16X32,
TxSize::Tx64x16 => return &DEFAULT_SCAN_32X16,
_ => {}
}
if tx_size.sqr_up_idx() == 4 {
return &DEFAULT_SCAN_32X32;
}
match tx_type {
TxType::Idtx => default_scan(tx_size),
TxType::VDct | TxType::VAdst | TxType::VFlipadst => mrow_scan(tx_size),
TxType::HDct | TxType::HAdst | TxType::HFlipadst => mcol_scan(tx_size),
_ => default_scan(tx_size),
}
}
fn default_scan(tx_size: TxSize) -> &'static [u16] {
match tx_size {
TxSize::Tx4x4 => &DEFAULT_SCAN_4X4,
TxSize::Tx4x8 => &DEFAULT_SCAN_4X8,
TxSize::Tx8x4 => &DEFAULT_SCAN_8X4,
TxSize::Tx8x8 => &DEFAULT_SCAN_8X8,
TxSize::Tx8x16 => &DEFAULT_SCAN_8X16,
TxSize::Tx16x8 => &DEFAULT_SCAN_16X8,
TxSize::Tx16x16 => &DEFAULT_SCAN_16X16,
TxSize::Tx16x32 => &DEFAULT_SCAN_16X32,
TxSize::Tx32x16 => &DEFAULT_SCAN_32X16,
TxSize::Tx4x16 => &DEFAULT_SCAN_4X16,
TxSize::Tx16x4 => &DEFAULT_SCAN_16X4,
TxSize::Tx8x32 => &DEFAULT_SCAN_8X32,
TxSize::Tx32x8 => &DEFAULT_SCAN_32X8,
_ => &DEFAULT_SCAN_32X32,
}
}
fn mrow_scan(tx_size: TxSize) -> &'static [u16] {
match tx_size {
TxSize::Tx4x4 => &MROW_SCAN_4X4,
TxSize::Tx4x8 => &MROW_SCAN_4X8,
TxSize::Tx8x4 => &MROW_SCAN_8X4,
TxSize::Tx8x8 => &MROW_SCAN_8X8,
TxSize::Tx8x16 => &MROW_SCAN_8X16,
TxSize::Tx16x8 => &MROW_SCAN_16X8,
TxSize::Tx16x16 => &MROW_SCAN_16X16,
TxSize::Tx4x16 => &MROW_SCAN_4X16,
_ => &MROW_SCAN_16X4,
}
}
fn mcol_scan(tx_size: TxSize) -> &'static [u16] {
match tx_size {
TxSize::Tx4x4 => &MCOL_SCAN_4X4,
TxSize::Tx4x8 => &MCOL_SCAN_4X8,
TxSize::Tx8x4 => &MCOL_SCAN_8X4,
TxSize::Tx8x8 => &MCOL_SCAN_8X8,
TxSize::Tx8x16 => &MCOL_SCAN_8X16,
TxSize::Tx16x8 => &MCOL_SCAN_16X8,
TxSize::Tx16x16 => &MCOL_SCAN_16X16,
TxSize::Tx4x16 => &MCOL_SCAN_4X16,
_ => &MCOL_SCAN_16X4,
}
}
pub fn decode_coeffs(
dec: &mut SymbolDecoder<'_>,
cdfs: &mut CoeffCdfs,
tx_size: TxSize,
tx: TxTypeCtx<'_>,
ptype: usize,
all_zero_ctx: usize,
dc_sign_ctx: usize,
) -> Result<CoeffBlock> {
let mut quant = [0_i32; MAX_COEFFS];
let pt = ptype.min(1);
let tx_ctx = tx_size.tx_size_ctx();
let skip_cdf = cdf_row(cdf_row(&mut cdfs.txb_skip, tx_ctx)?, all_zero_ctx)?;
let all_zero = dec.read_symbol(skip_cdf)? != 0;
if all_zero {
return Ok(CoeffBlock {
quant,
eob: 0,
cul_level: 0,
dc_category: 0,
tx_type: TxType::DctDct,
});
}
let tx_type = if ptype == 0 {
read_transform_type(
dec,
tx.intra_cdfs,
tx.set,
tx_size,
tx.intra_dir,
tx.qindex_positive,
)?
} else if !tx.lossless {
chroma_tx_type(tx.uv_mode, tx.set)
} else {
TxType::DctDct
};
let cls = tx_class(tx_type);
let scan = get_scan(tx_size, tx_type);
let eob_ctx = usize::from(cls != 0);
let eob_pt = match tx_size.eob_multisize() {
0 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_16, pt)?, eob_ctx)?)?,
1 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_32, pt)?, eob_ctx)?)?,
2 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_64, pt)?, eob_ctx)?)?,
3 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_128, pt)?, eob_ctx)?)?,
4 => dec.read_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_256, pt)?, eob_ctx)?)?,
5 => dec.read_symbol(cdf_row(&mut cdfs.eob_pt_512, pt)?)?,
_ => dec.read_symbol(cdf_row(&mut cdfs.eob_pt_1024, pt)?)?,
} + 1;
let mut eob = if eob_pt < 2 {
eob_pt
} else {
(1 << (eob_pt - 2)) + 1
};
if let Some(eob_shift) = eob_pt.checked_sub(3) {
let extra_cdf = cdf_row(
cdf_row(cdf_row(&mut cdfs.eob_extra, tx_ctx)?, pt)?,
eob_pt - 3,
)?;
if dec.read_symbol(extra_cdf)? != 0 {
eob += 1 << eob_shift;
}
for i in 1..eob_pt.saturating_sub(2) {
let shift = eob_pt.saturating_sub(2) - 1 - i;
if dec.read_bool()? {
eob += 1 << shift;
}
}
}
eob = eob.min(scan.len());
for c in (0..eob).rev() {
let pos = scan.get(c).map_or(0, |&p| usize::from(p));
let mut level;
if c == eob - 1 {
let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, true) + SIG_COEF_CONTEXTS_EOB
- SIG_COEF_CONTEXTS;
let cdf_ref = cdf_row(
cdf_row(cdf_row(&mut cdfs.coeff_base_eob, tx_ctx)?, pt)?,
ctx,
)?;
level = dec.read_symbol(cdf_ref)? as i32 + 1;
} else {
let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, false);
let cdf_ref = cdf_row(cdf_row(cdf_row(&mut cdfs.coeff_base, tx_ctx)?, pt)?, ctx)?;
level = dec.read_symbol(cdf_ref)? as i32;
}
if level > NUM_BASE_LEVELS {
let br_ctx = coeff_br_ctx(tx_size, cls, &quant, pos);
let br_bucket = tx_ctx.min(3);
for _ in 0..(COEFF_BASE_RANGE / (BR_CDF_SIZE - 1)) {
let cdf_ref = cdf_row(
cdf_row(cdf_row(&mut cdfs.coeff_br, br_bucket)?, pt)?,
br_ctx,
)?;
let coeff_br = dec.read_symbol(cdf_ref)? as i32;
level += coeff_br;
if coeff_br < BR_CDF_SIZE - 1 {
break;
}
}
}
if let Some(slot) = quant.get_mut(pos) {
*slot = level;
}
}
let mut cul_level: i32 = 0;
let mut dc_category = 0_u8;
for c in 0..eob {
let pos = scan.get(c).map_or(0, |&p| usize::from(p));
let level = quant.get(pos).copied().unwrap_or(0);
let sign = if level != 0 {
if c == 0 {
let cdf_ref = cdf_row(cdf_row(&mut cdfs.dc_sign, pt)?, dc_sign_ctx)?;
dec.read_symbol(cdf_ref)? != 0
} else {
dec.read_bool()?
}
} else {
false
};
let mut magnitude = level;
if magnitude > NUM_BASE_LEVELS + COEFF_BASE_RANGE {
magnitude = read_golomb(dec)? + COEFF_BASE_RANGE + NUM_BASE_LEVELS;
}
if pos == 0 && magnitude > 0 {
dc_category = if sign { 1 } else { 2 };
}
magnitude &= 0xF_FFFF;
cul_level += magnitude;
if let Some(slot) = quant.get_mut(pos) {
*slot = if sign { -magnitude } else { magnitude };
}
}
Ok(CoeffBlock {
quant,
eob,
cul_level: cul_level.min(63) as u8,
dc_category,
tx_type,
})
}
#[allow(clippy::too_many_arguments, reason = "the coefficient syntax's inputs")]
pub(crate) fn encode_coeffs(
enc: &mut SymbolEncoder,
cdfs: &mut CoeffCdfs,
tx_size: TxSize,
tx: TxTypeCtx<'_>,
tx_type: TxType,
ptype: usize,
all_zero_ctx: usize,
dc_sign_ctx: usize,
levels: &[i32],
) -> Result<CoeffBlock> {
let pt = ptype.min(1);
let tx_ctx = tx_size.tx_size_ctx();
let cls = tx_class(tx_type);
let scan = get_scan(tx_size, tx_type);
let level_at = |c: usize| {
levels
.get(scan.get(c).map_or(0, |&p| usize::from(p)))
.copied()
.unwrap_or(0)
};
let eob = (0..scan.len())
.rev()
.find(|&c| level_at(c) != 0)
.map_or(0, |c| c + 1);
let skip_cdf = cdf_row(cdf_row(&mut cdfs.txb_skip, tx_ctx)?, all_zero_ctx)?;
enc.write_symbol(skip_cdf, usize::from(eob == 0));
let mut quant = [0_i32; MAX_COEFFS];
if eob == 0 {
return Ok(CoeffBlock {
quant,
eob: 0,
cul_level: 0,
dc_category: 0,
tx_type: TxType::DctDct,
});
}
if ptype == 0 {
write_transform_type(
enc,
tx.intra_cdfs,
tx.set,
tx_size,
tx.intra_dir,
tx.qindex_positive,
tx_type,
)?;
}
let eob_pt = if eob <= 2 {
eob
} else {
2 + (usize::BITS - 1 - (eob - 1).leading_zeros()) as usize
};
let eob_ctx = usize::from(cls != 0);
let symbol = eob_pt - 1;
match tx_size.eob_multisize() {
0 => enc.write_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_16, pt)?, eob_ctx)?, symbol),
1 => enc.write_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_32, pt)?, eob_ctx)?, symbol),
2 => enc.write_symbol(cdf_row(cdf_row(&mut cdfs.eob_pt_64, pt)?, eob_ctx)?, symbol),
3 => enc.write_symbol(
cdf_row(cdf_row(&mut cdfs.eob_pt_128, pt)?, eob_ctx)?,
symbol,
),
4 => enc.write_symbol(
cdf_row(cdf_row(&mut cdfs.eob_pt_256, pt)?, eob_ctx)?,
symbol,
),
5 => enc.write_symbol(cdf_row(&mut cdfs.eob_pt_512, pt)?, symbol),
_ => enc.write_symbol(cdf_row(&mut cdfs.eob_pt_1024, pt)?, symbol),
}
if eob_pt >= 3 {
let offset = eob - ((1 << (eob_pt - 2)) + 1);
let extra_cdf = cdf_row(
cdf_row(cdf_row(&mut cdfs.eob_extra, tx_ctx)?, pt)?,
eob_pt - 3,
)?;
enc.write_symbol(extra_cdf, (offset >> (eob_pt - 3)) & 1);
for i in 1..eob_pt - 2 {
let shift = eob_pt - 2 - 1 - i;
enc.write_bool((offset >> shift) & 1 == 1);
}
}
let cap = NUM_BASE_LEVELS + COEFF_BASE_RANGE + 1;
for c in (0..eob).rev() {
let pos = scan.get(c).map_or(0, |&p| usize::from(p));
let magnitude = level_at(c).abs().min(cap);
if c == eob - 1 {
let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, true) + SIG_COEF_CONTEXTS_EOB
- SIG_COEF_CONTEXTS;
let cdf_ref = cdf_row(
cdf_row(cdf_row(&mut cdfs.coeff_base_eob, tx_ctx)?, pt)?,
ctx,
)?;
enc.write_symbol(cdf_ref, (magnitude.min(3) - 1) as usize);
} else {
let ctx = coeff_base_ctx(tx_size, cls, &quant, pos, c, false);
let cdf_ref = cdf_row(cdf_row(cdf_row(&mut cdfs.coeff_base, tx_ctx)?, pt)?, ctx)?;
enc.write_symbol(cdf_ref, magnitude.min(3) as usize);
}
if magnitude > NUM_BASE_LEVELS {
let br_ctx = coeff_br_ctx(tx_size, cls, &quant, pos);
let br_bucket = tx_ctx.min(3);
let mut remaining = magnitude - NUM_BASE_LEVELS - 1;
for _ in 0..(COEFF_BASE_RANGE / (BR_CDF_SIZE - 1)) {
let k = remaining.min(BR_CDF_SIZE - 1);
let cdf_ref = cdf_row(
cdf_row(cdf_row(&mut cdfs.coeff_br, br_bucket)?, pt)?,
br_ctx,
)?;
enc.write_symbol(cdf_ref, k as usize);
remaining -= k;
if k < BR_CDF_SIZE - 1 {
break;
}
}
}
if let Some(slot) = quant.get_mut(pos) {
*slot = magnitude;
}
}
let mut cul_level = 0_i32;
let mut dc_category = 0_u8;
for c in 0..eob {
let pos = scan.get(c).map_or(0, |&p| usize::from(p));
let value = level_at(c);
let magnitude = value.abs();
if magnitude != 0 {
if c == 0 {
let cdf_ref = cdf_row(cdf_row(&mut cdfs.dc_sign, pt)?, dc_sign_ctx)?;
enc.write_symbol(cdf_ref, usize::from(value < 0));
} else {
enc.write_bool(value < 0);
}
}
if magnitude > NUM_BASE_LEVELS + COEFF_BASE_RANGE {
write_golomb(enc, (magnitude - NUM_BASE_LEVELS - COEFF_BASE_RANGE) as u32);
}
if pos == 0 && magnitude > 0 {
dc_category = if value < 0 { 1 } else { 2 };
}
let magnitude = magnitude & 0xF_FFFF;
cul_level += magnitude;
if let Some(slot) = quant.get_mut(pos) {
*slot = if value < 0 { -magnitude } else { magnitude };
}
}
Ok(CoeffBlock {
quant,
eob,
cul_level: cul_level.min(63) as u8,
dc_category,
tx_type,
})
}
fn write_golomb(enc: &mut SymbolEncoder, x: u32) {
let length = 32 - x.leading_zeros();
for _ in 1..length {
enc.write_bool(false);
}
enc.write_bool(true);
for i in (0..length - 1).rev() {
enc.write_bool((x >> i) & 1 == 1);
}
}
fn coeff_base_ctx(
tx_size: TxSize,
cls: usize,
quant: &[i32; MAX_COEFFS],
pos: usize,
c: usize,
is_eob: bool,
) -> usize {
let bwl = tx_size.adjusted_log2_width();
let width = tx_size.adjusted_width() as i32;
let height = tx_size.adjusted_height() as i32;
if is_eob {
let area = (height as usize) << bwl;
return if c == 0 {
SIG_COEF_CONTEXTS - 4
} else if c <= area / 8 {
SIG_COEF_CONTEXTS - 3
} else if c <= area / 4 {
SIG_COEF_CONTEXTS - 2
} else {
SIG_COEF_CONTEXTS - 1
};
}
let row = (pos >> bwl) as i32;
let col = (pos - ((row as usize) << bwl)) as i32;
let mut mag = 0;
for &(d_row, d_col) in offsets_sig(cls) {
let ref_row = row + d_row;
let ref_col = col + d_col;
if ref_row >= 0 && ref_col >= 0 && ref_row < height && ref_col < width {
let ref_pos = ((ref_row as usize) << bwl) + ref_col as usize;
mag += quant.get(ref_pos).copied().unwrap_or(0).abs().min(3);
}
}
let ctx = ((mag + 1) >> 1).min(4);
if cls == 0 {
if row == 0 && col == 0 {
return 0;
}
let offset = COEFF_BASE_CTX_OFFSET
.get(tx_size as usize)
.and_then(|t| t.get(row.min(4) as usize))
.and_then(|r| r.get(col.min(4) as usize))
.copied()
.unwrap_or(0);
return (ctx + offset) as usize;
}
let idx = if cls == 2 { row } else { col };
let offset = pick3(COEFF_BASE_POS_CTX_OFFSET, idx.min(2) as usize);
(ctx + offset) as usize
}
fn coeff_br_ctx(tx_size: TxSize, cls: usize, quant: &[i32; MAX_COEFFS], pos: usize) -> usize {
let bwl = tx_size.adjusted_log2_width();
let txw = tx_size.adjusted_width();
let txh = tx_size.adjusted_height() as i32;
let row = (pos >> bwl) as i32;
let col = (pos - ((row as usize) << bwl)) as i32;
let mut mag = 0;
for &(d_row, d_col) in offsets_mag(cls) {
let ref_row = row + d_row;
let ref_col = col + d_col;
if ref_row >= 0 && ref_col >= 0 && ref_row < txh && ref_col < (1 << bwl) {
let ref_pos = ref_row as usize * txw + ref_col as usize;
mag += quant
.get(ref_pos)
.copied()
.unwrap_or(0)
.min(COEFF_BASE_RANGE + NUM_BASE_LEVELS + 1);
}
}
let mag = ((mag + 1) >> 1).min(6);
let ctx = if pos == 0 {
mag
} else if cls == 0 {
if row < 2 && col < 2 {
mag + 7
} else {
mag + 14
}
} else if cls == 1 {
if col == 0 { mag + 7 } else { mag + 14 }
} else if row == 0 {
mag + 7
} else {
mag + 14
};
ctx as usize
}
fn offsets_sig(cls: usize) -> &'static [(i32, i32); 5] {
let [two_d, horiz, vert] = &SIG_REF_DIFF_OFFSET;
match cls {
1 => horiz,
2 => vert,
_ => two_d,
}
}
fn offsets_mag(cls: usize) -> &'static [(i32, i32); 3] {
let [two_d, horiz, vert] = &MAG_REF_OFFSET;
match cls {
1 => horiz,
2 => vert,
_ => two_d,
}
}
fn pick3<T: Copy>(arr: [T; 3], q: usize) -> T {
let [a, b, c] = arr;
match q {
1 => b,
2 => c,
_ => a,
}
}
fn read_golomb(dec: &mut SymbolDecoder<'_>) -> Result<i32> {
let mut length = 0_i32;
loop {
length += 1;
if dec.read_bool()? {
break;
}
if length > 20 {
break;
}
}
let mut x = 1_i32;
for _ in 0..length.saturating_sub(1) {
x = (x << 1) | i32::from(dec.read_bool()?);
}
Ok(x)
}
#[cfg(test)]
#[allow(
clippy::unwrap_used,
clippy::indexing_slicing,
clippy::panic,
reason = "tests operate on known-good values and assert shapes directly"
)]
mod tests {
use super::*;
#[test]
fn written_coefficients_decode_to_the_same_block() {
use super::super::transform_type::intra_tx_set;
let mut state = 0x9e37_79b9_u32;
let mut next = move || {
state ^= state << 13;
state ^= state >> 17;
state ^= state << 5;
state
};
let sizes = [
TxSize::Tx4x4,
TxSize::Tx8x8,
TxSize::Tx16x16,
TxSize::Tx32x32,
TxSize::Tx8x16,
TxSize::Tx16x4,
];
for round in 0..60 {
let size = sizes[round % sizes.len()];
let ptype = usize::from(round % 3 == 2);
let set = intra_tx_set(size, false);
let uv_mode = 1; let tx_type = if ptype == 0 {
TxType::DctDct
} else {
chroma_tx_type(uv_mode, set)
};
let (w, h) = (size.adjusted_width(), size.adjusted_height());
let density = next() % 4;
let levels: Vec<i32> = (0..w * h)
.map(|_| {
let r = next();
if density == 0 || r % (2 + density * 3) != 0 {
0
} else {
let m = match r % 7 {
0 => 40 + (r >> 20) as i32 % 3000,
1 => 3 + (r >> 8) as i32 % 12,
_ => 1 + (r >> 8) as i32 % 2,
};
if r & 1 == 0 { m } else { -m }
}
})
.collect();
let (az, ds) = ((next() % 13) as usize, (next() % 3) as usize);
let mut enc = SymbolEncoder::new(false);
let mut cdfs = CoeffCdfs::new(2);
let mut tt = IntraTxTypeCdfs::new();
let written = encode_coeffs(
&mut enc,
&mut cdfs,
size,
TxTypeCtx {
set,
intra_cdfs: &mut tt,
intra_dir: 0,
uv_mode,
qindex_positive: true,
lossless: false,
},
tx_type,
ptype,
az,
ds,
&levels,
)
.unwrap();
let data = enc.finish();
let mut dec = SymbolDecoder::new(&data, false).unwrap();
let mut cdfs = CoeffCdfs::new(2);
let mut tt = IntraTxTypeCdfs::new();
let read = decode_coeffs(
&mut dec,
&mut cdfs,
size,
TxTypeCtx {
set,
intra_cdfs: &mut tt,
intra_dir: 0,
uv_mode,
qindex_positive: true,
lossless: false,
},
ptype,
az,
ds,
)
.unwrap();
assert_eq!(read.eob, written.eob, "round {round}");
assert_eq!(read.quant[..w * h], written.quant[..w * h], "round {round}");
assert_eq!(read.quant[..w * h], levels[..], "round {round}: levels");
assert_eq!(
(read.cul_level, read.dc_category, read.tx_type),
(written.cul_level, written.dc_category, written.tx_type)
);
}
}
#[test]
fn base_eob_contexts_match_the_spec_buckets() {
let q = [0; MAX_COEFFS];
assert_eq!(
coeff_base_ctx(TxSize::Tx4x4, 0, &q, 0, 0, true),
SIG_COEF_CONTEXTS - 4
);
assert_eq!(
coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 1, true),
SIG_COEF_CONTEXTS - 3
);
assert_eq!(
coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 3, true),
SIG_COEF_CONTEXTS - 2
);
assert_eq!(
coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 9, true),
SIG_COEF_CONTEXTS - 1
);
}
#[test]
fn dc_position_base_context_is_zero() {
let q = [3; MAX_COEFFS];
assert_eq!(coeff_base_ctx(TxSize::Tx4x4, 0, &q, 0, 4, false), 0);
}
#[test]
fn base_context_folds_in_neighbour_magnitudes() {
let mut q = [0_i32; MAX_COEFFS];
q[6] = 3;
assert_eq!(coeff_base_ctx(TxSize::Tx4x4, 0, &q, 5, 4, false), 8);
}
#[test]
fn br_context_at_dc_is_the_bare_magnitude() {
let mut q = [0_i32; MAX_COEFFS];
q[1] = 5;
assert_eq!(coeff_br_ctx(TxSize::Tx4x4, 0, &q, 0), 3);
}
#[test]
fn vertical_class_uses_position_offsets() {
let q = [0_i32; MAX_COEFFS];
assert_eq!(
coeff_base_ctx(TxSize::Tx8x8, 2, &q, 8, 4, false),
(SIG_COEF_CONTEXTS_2D + 5) as usize
);
}
#[test]
fn scans_are_selected_by_size_and_type() {
assert_eq!(get_scan(TxSize::Tx4x4, TxType::DctDct).len(), 16);
assert_eq!(get_scan(TxSize::Tx8x8, TxType::DctDct).len(), 64);
assert_eq!(get_scan(TxSize::Tx64x64, TxType::DctDct).len(), 1024);
assert_eq!(get_scan(TxSize::Tx16x64, TxType::DctDct).len(), 512);
assert_eq!(get_scan(TxSize::Tx8x8, TxType::VDct), &MROW_SCAN_8X8[..]);
assert_eq!(get_scan(TxSize::Tx8x8, TxType::HDct), &MCOL_SCAN_8X8[..]);
assert_ne!(
get_scan(TxSize::Tx8x8, TxType::DctDct),
get_scan(TxSize::Tx8x8, TxType::VDct)
);
}
#[test]
fn an_all_zero_block_reads_one_symbol_and_stops() {
let data = [0x00; 8];
let mut dec = SymbolDecoder::new(&data, true).unwrap();
let mut cdfs = CoeffCdfs::new(0);
let mut intra = IntraTxTypeCdfs::new();
let tx = TxTypeCtx {
set: IntraTxSet::Set1,
intra_cdfs: &mut intra,
intra_dir: 0,
uv_mode: 0,
qindex_positive: false,
lossless: true,
};
let block = decode_coeffs(&mut dec, &mut cdfs, TxSize::Tx4x4, tx, 0, 0, 0).unwrap();
assert_eq!(block.tx_type, TxType::DctDct);
if block.eob == 0 {
assert_eq!(block.quant, [0; MAX_COEFFS]);
assert_eq!(block.cul_level, 0);
}
}
#[test]
fn golomb_reads_a_unary_prefix_then_data_bits() {
let data = [0xFF; 4];
let mut dec = SymbolDecoder::new(&data, true).unwrap();
assert_eq!(read_golomb(&mut dec).unwrap(), 1);
}
}