#![allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::cast_sign_loss,
reason = "quantized levels (0..=2047) and reconstructed coefficients are stored \
into i16 with the reference decoder's int16_t wrapping semantics, and \
the 0..=15 zig-zag positions cast to i32; every value is in range and \
the casts reproduce the decoder's arithmetic exactly"
)]
use crate::lossy::constants::{BANDS, CAT_3456, CoeffProbas, NUM_PROBAS, Prob, ZIGZAG};
use crate::lossy::prob_opt::bit_cost;
use crate::lossy::quant::{QPair, Quantized};
use crate::lossy::work::work;
const MAX_LEVEL: i32 = 2047;
pub(crate) const RD_DISTO_MULT: i64 = 256;
const INF: i64 = i64::MAX / 4;
const LAMBDA_SHIFT: u32 = 6;
pub(crate) const fn trellis_lambda(q_ac: i32) -> i64 {
let q = q_ac as i64;
let l = (7 * q * q) >> LAMBDA_SHIFT;
if l < 1 { 1 } else { l }
}
const fn ctx_of(m: i32) -> usize {
if m == 0 {
0
} else if m == 1 {
1
} else {
2
}
}
pub(crate) fn large_value_rate(m: i32, p: [Prob; NUM_PROBAS]) -> i64 {
if m <= 4 {
let mut c = i64::from(bit_cost(false, p[3]));
if m == 2 {
c += i64::from(bit_cost(false, p[4]));
} else {
c += i64::from(bit_cost(true, p[4]));
c += i64::from(bit_cost(m == 4, p[5]));
}
c
} else if m <= 10 {
let mut c = i64::from(bit_cost(true, p[3])) + i64::from(bit_cost(false, p[6]));
if m <= 6 {
c += i64::from(bit_cost(false, p[7]));
c += i64::from(bit_cost(m == 6, 159));
} else {
c += i64::from(bit_cost(true, p[7]));
let hi = m - 7; c += i64::from(bit_cost((hi >> 1) & 1 == 1, 165));
c += i64::from(bit_cost(hi & 1 == 1, 145));
}
c
} else {
let mut c = i64::from(bit_cost(true, p[3])) + i64::from(bit_cost(true, p[6]));
let cat = match m {
11..=18 => 0usize,
19..=34 => 1,
35..=66 => 2,
_ => 3,
};
c += i64::from(bit_cost((cat >> 1) & 1 == 1, p[8]));
c += i64::from(bit_cost(cat & 1 == 1, p[9 + (cat >> 1)]));
let extra = m - 3 - (8 << cat);
let probs = CAT_3456[cat];
let nbits = probs.len();
for (i, &prob) in probs.iter().enumerate() {
c += i64::from(bit_cost((extra >> (nbits - 1 - i)) & 1 == 1, prob));
}
c
}
}
fn value_rate(m: i32, p: [Prob; NUM_PROBAS]) -> i64 {
if m == 0 {
return i64::from(bit_cost(false, p[1]));
}
let mut c = i64::from(bit_cost(true, p[1])) + 256; if m == 1 {
c += i64::from(bit_cost(false, p[2]));
} else {
c += i64::from(bit_cost(true, p[2]));
c += large_value_rate(m, p);
}
c
}
#[must_use]
pub(crate) fn block_token_cost(
levels: [i16; 16],
first: usize,
last: i32,
plane: usize,
ctx0: usize,
bands: &CoeffProbas,
) -> i64 {
work!(TokenCostWalk);
let pb = &bands[plane];
let mut n = first;
let mut band = BANDS[n];
let mut ctx = ctx0;
let mut cost = 0i64;
loop {
let p = pb[band][ctx];
if n as i32 > last {
return cost + i64::from(bit_cost(false, p[0])); }
cost += i64::from(bit_cost(true, p[0])); while levels[n] == 0 {
let pz = pb[band][ctx];
cost += i64::from(bit_cost(false, pz[1]));
n += 1;
band = BANDS[n];
ctx = 0;
}
let p = pb[band][ctx];
cost += i64::from(bit_cost(true, p[1])); let v = i32::from(levels[n]).abs();
if v == 1 {
cost += i64::from(bit_cost(false, p[2]));
ctx = 1;
} else {
cost += i64::from(bit_cost(true, p[2]));
cost += large_value_rate(v, p);
ctx = 2;
}
cost += 256; n += 1;
if n == 16 {
return cost; }
band = BANDS[n];
}
}
struct Viterbi {
cost: [[i64; 3]; 16],
mag: [[i32; 3]; 16],
prev: [[usize; 3]; 16],
}
impl Viterbi {
const fn backtrack(
&self,
coeffs: [i16; 16],
pair: QPair,
first: usize,
best_n: usize,
best_out: usize,
) -> Quantized {
let mut levels = [0i16; 16];
let mut recon = [0i16; 16];
let (mut out, mut n) = (best_out, best_n);
loop {
let mag = self.mag[n][out];
let j = ZIGZAG[n];
let q = if j == 0 { pair.dc.q } else { pair.ac.q };
let signed = if coeffs[j] < 0 { -mag } else { mag };
levels[n] = signed as i16;
recon[j] = (signed * q) as i16;
if n == first {
break;
}
out = self.prev[n][out];
n -= 1;
}
Quantized {
levels,
recon,
last: best_n as i32,
}
}
}
#[must_use]
pub(crate) fn trellis_quantize_block(
coeffs: [i16; 16],
pair: QPair,
first: usize,
ctx0: usize,
plane: usize,
bands: &CoeffProbas,
lambda: i64,
) -> Quantized {
let pb = &bands[plane];
let mut suffix_dist = [0i64; 17];
for n in (first..16).rev() {
let c = i64::from(coeffs[ZIGZAG[n]]);
suffix_dist[n] = suffix_dist[n + 1] + c * c;
}
let mut v = Viterbi {
cost: [[INF; 3]; 16],
mag: [[0i32; 3]; 16],
prev: [[0usize; 3]; 16],
};
let mut prev_cost = [INF; 3];
prev_cost[ctx0] = 0;
let first_extra = if ctx0 == 0 {
i64::from(bit_cost(true, pb[BANDS[first]][0][0]))
} else {
0
};
let (mut best_total, mut best_n, mut best_out) = (INF, 0usize, 0usize);
for n in first..16 {
let j = ZIGZAG[n];
let factor = if j == 0 { pair.dc } else { pair.ac };
let (q, abs_coeff) = (factor.q, i32::from(coeffs[j]).abs());
let l0 = factor.quantize(i32::from(coeffs[j])).abs(); debug_assert!(
l0 <= MAX_LEVEL,
"quantize clamps the nearest level to MAX_LEVEL"
);
let mut cands = [l0, 0, 0];
let mut ncand = 1usize;
if l0 >= 1 {
cands[ncand] = l0 - 1;
ncand += 1;
}
if l0 >= 2 {
cands[ncand] = 0;
ncand += 1;
}
let cands = &cands[..ncand];
for ci in 0..3 {
if prev_cost[ci] >= INF {
continue;
}
let p = pb[BANDS[n]][ci];
let base = if ci > 0 {
i64::from(bit_cost(true, p[0]))
} else {
0
};
let extra = if n == first { first_extra } else { 0 };
for &m in cands {
work!(TrellisEval);
let err = i64::from(abs_coeff - m * q);
let rate = base + extra + value_rate(m, p);
let cand = prev_cost[ci] + RD_DISTO_MULT * err * err + lambda * rate;
let out = ctx_of(m);
if cand < v.cost[n][out] {
v.cost[n][out] = cand;
v.mag[n][out] = m;
v.prev[n][out] = ci;
}
}
}
for out in [1usize, 2] {
if v.cost[n][out] >= INF {
continue;
}
let eob = if n == 15 {
0
} else {
i64::from(bit_cost(false, pb[BANDS[n + 1]][out][0]))
};
let total = v.cost[n][out] + lambda * eob + RD_DISTO_MULT * suffix_dist[n + 1];
if total < best_total {
(best_total, best_n, best_out) = (total, n, out);
}
}
prev_cost = v.cost[n];
}
let empty_total = RD_DISTO_MULT * suffix_dist[first]
+ lambda * i64::from(bit_cost(false, pb[BANDS[first]][ctx0][0]));
if best_total >= empty_total {
return Quantized {
levels: [0; 16],
recon: [0; 16],
last: first as i32 - 1,
};
}
v.backtrack(coeffs, pair, first, best_n, best_out)
}
#[cfg(test)]
mod tests {
use super::{block_token_cost, trellis_lambda, trellis_quantize_block};
use crate::lossy::constants::{BANDS, COEFFS_PROBA_0, NUM_PROBAS, Prob, ZIGZAG};
use crate::lossy::prob_opt::bit_cost;
use crate::lossy::quant::{Quantizer, quantize_block};
#[test]
fn block_token_cost_matches_the_hand_walked_token_tree() {
let (plane, first, ctx0) = (0usize, 1usize, 0usize);
let empty = [0i16; 16];
let got = block_token_cost(empty, first, first as i32 - 1, plane, ctx0, &COEFFS_PROBA_0);
let p = COEFFS_PROBA_0[plane][BANDS[first]][ctx0];
assert_eq!(
got,
i64::from(bit_cost(false, p[0])),
"empty block = one EOB bit"
);
let mut levels = [0i16; 16];
levels[first] = 1;
let got = block_token_cost(levels, first, first as i32, plane, ctx0, &COEFFS_PROBA_0);
let p0 = COEFFS_PROBA_0[plane][BANDS[first]][ctx0];
let eob = COEFFS_PROBA_0[plane][BANDS[first + 1]][1];
let want = i64::from(bit_cost(true, p0[0]))
+ i64::from(bit_cost(true, p0[1]))
+ i64::from(bit_cost(false, p0[2]))
+ 256
+ i64::from(bit_cost(false, eob[0]));
assert_eq!(got, want, "one-unit block = more+nonzero+one+sign+EOB");
}
#[test]
fn block_token_cost_charges_more_for_a_busier_block() {
let mut sparse = [0i16; 16];
sparse[0] = 1;
let mut busy = [0i16; 16];
busy[0] = 5;
busy[1] = -3;
busy[2] = 17;
busy[5] = 2;
let c_sparse = block_token_cost(sparse, 0, 0, 0, 0, &COEFFS_PROBA_0);
let c_busy = block_token_cost(busy, 0, 5, 0, 0, &COEFFS_PROBA_0);
assert!(
c_busy > c_sparse,
"busy {c_busy} should cost more than sparse {c_sparse}"
);
}
fn assert_valid(q: Quantizer, first: usize, quantized: &super::Quantized) {
let (dc, ac) = (q.y1.dc, q.y1.ac);
let mut want_last = first as i32 - 1;
for (n, &j) in ZIGZAG.iter().enumerate().skip(first) {
let level = i32::from(quantized.levels[n]);
assert!(level.abs() <= 2047, "level {level} at {n} out of range");
let step = if j == 0 { dc.q } else { ac.q };
assert_eq!(
i32::from(quantized.recon[j]),
level * step,
"recon[{j}] must equal level*q"
);
if level != 0 {
want_last = n as i32;
}
}
assert_eq!(
quantized.last, want_last,
"last must be the last non-zero index"
);
}
#[test]
fn trellis_result_is_valid_and_self_consistent_in_shape() {
let q = Quantizer::new(40);
let mut coeffs = [0i16; 16];
coeffs[0] = 90;
coeffs[1] = -50;
coeffs[4] = 12;
coeffs[8] = 200;
let lambda = trellis_lambda(q.y1.ac.q);
let out = trellis_quantize_block(coeffs, q.y1, 0, 0, 0, &COEFFS_PROBA_0, lambda);
assert_valid(q, 0, &out);
}
#[test]
fn tiny_coefficients_trellis_toward_zero() {
let q = Quantizer::new(80); let mut coeffs = [0i16; 16];
coeffs[1] = 3;
coeffs[4] = -2;
coeffs[5] = 4;
let lambda = trellis_lambda(q.y1.ac.q);
let out = trellis_quantize_block(coeffs, q.y1, 1, 0, 0, &COEFFS_PROBA_0, lambda);
assert_valid(q, 1, &out);
assert_eq!(
out.last, 0,
"tiny AC-only block should trellis to empty (last = first-1)"
);
}
#[test]
fn trellis_never_scores_worse_than_round_to_nearest_and_is_deterministic() {
let q = Quantizer::new(60);
let mut coeffs = [0i16; 16];
coeffs[0] = 140;
coeffs[1] = 41;
coeffs[2] = -39;
coeffs[8] = 8;
let rn = quantize_block(coeffs, q.y1.dc, q.y1.ac, 0);
let lambda = trellis_lambda(q.y1.ac.q);
let a = trellis_quantize_block(coeffs, q.y1, 0, 0, 0, &COEFFS_PROBA_0, lambda);
let b = trellis_quantize_block(coeffs, q.y1, 0, 0, 0, &COEFFS_PROBA_0, lambda);
assert_eq!(a.levels, b.levels, "trellis must be deterministic");
assert_eq!(a.recon, b.recon);
assert!(
a.last <= rn.last,
"trellis last {} > round-to-nearest {}",
a.last,
rn.last
);
assert_valid(q, 0, &a);
}
#[test]
fn lambda_grows_with_the_quantizer_step() {
assert_eq!(trellis_lambda(4).max(1), trellis_lambda(4));
assert!(trellis_lambda(200) > trellis_lambda(20));
assert!(trellis_lambda(1) >= 1, "lambda is clamped to at least 1");
}
#[test]
fn trellis_lambda_is_exact() {
assert_eq!(trellis_lambda(1), 1);
assert_eq!(trellis_lambda(4), 1);
assert_eq!(trellis_lambda(20), 43);
assert_eq!(trellis_lambda(100), 1093);
assert_eq!(trellis_lambda(200), 4375);
}
const RATE_P: [Prob; NUM_PROBAS] = [128, 128, 128, 40, 50, 60, 70, 80, 90, 100, 110];
#[test]
fn large_value_rate_is_exact_across_every_category() {
use super::large_value_rate;
let want = [
(2, 1289),
(3, 1302),
(4, 865),
(5, 1148),
(6, 1330),
(7, 1052),
(8, 1151),
(9, 1272),
(10, 1371),
(11, 1484),
(12, 1553),
(18, 1941),
(19, 1532),
(25, 1759),
(34, 2092),
(35, 1673),
(50, 1966),
(66, 2285),
(67, 2006),
(100, 2310),
(300, 3621),
(700, 4751),
(2047, 8013),
];
for (m, w) in want {
assert_eq!(large_value_rate(m, RATE_P), w, "large_value_rate({m})");
}
}
#[test]
fn value_rate_is_exact() {
use super::value_rate;
let want = [
(0, 256),
(1, 768),
(2, 2057),
(3, 2070),
(5, 1916),
(8, 1919),
(20, 2341),
(100, 3078),
(2047, 8781),
];
for (m, w) in want {
assert_eq!(value_rate(m, RATE_P), w, "value_rate({m})");
}
}
#[test]
fn block_token_cost_exact_for_zero_run_and_large_value() {
let mut zero_run = [0i16; 16];
zero_run[2] = 3;
assert_eq!(
block_token_cost(zero_run, 0, 2, 0, 0, &COEFFS_PROBA_0),
3893
);
let mut mixed = [0i16; 16];
mixed[0] = 1;
mixed[3] = -2;
assert_eq!(block_token_cost(mixed, 0, 3, 0, 0, &COEFFS_PROBA_0), 4636);
let mut large = [0i16; 16];
large[1] = 4;
assert_eq!(block_token_cost(large, 1, 1, 0, 0, &COEFFS_PROBA_0), 5941);
}
fn block(pairs: &[(usize, i16)]) -> [i16; 16] {
let mut c = [0i16; 16];
for &(i, v) in pairs {
c[i] = v;
}
c
}
fn trellis(base_q: i32, coeffs: [i16; 16], first: usize, ctx0: usize) -> super::Quantized {
let q = Quantizer::new(base_q);
let lambda = trellis_lambda(q.y1.ac.q);
trellis_quantize_block(coeffs, q.y1, first, ctx0, 0, &COEFFS_PROBA_0, lambda)
}
#[test]
fn trellis_quantize_block_exact_goldens() {
struct Case {
q: i32,
coeffs: [i16; 16],
first: usize,
ctx0: usize,
levels: [i16; 16],
recon: [i16; 16],
last: i32,
}
let cases = [
Case {
q: 40,
coeffs: block(&[(0, 90), (1, -50), (4, 12), (8, 200)]),
first: 0,
ctx0: 0,
levels: block(&[(0, 2), (1, -1), (3, 5)]),
recon: block(&[(0, 74), (1, -44), (8, 220)]),
last: 3,
},
Case {
q: 20,
coeffs: block(&[(0, 100), (5, 60), (15, 400)]),
first: 0,
ctx0: 0,
levels: block(&[(0, 5), (4, 2), (15, 17)]),
recon: block(&[(0, 105), (5, 48), (15, 408)]),
last: 15,
},
Case {
q: 30,
coeffs: block(&[(0, 200), (8, 150)]),
first: 0,
ctx0: 0,
levels: block(&[(0, 7), (3, 4)]),
recon: block(&[(0, 189), (8, 136)]),
last: 3,
},
Case {
q: 40,
coeffs: block(&[(1, 90), (4, -50), (8, 30)]),
first: 1,
ctx0: 2,
levels: block(&[(1, 2), (2, -1)]),
recon: block(&[(1, 88), (4, -44)]),
last: 2,
},
Case {
q: 40,
coeffs: block(&[(0, 60), (1, 30), (2, -45), (5, 12)]),
first: 0,
ctx0: 1,
levels: block(&[(0, 2), (1, 1), (5, -1)]),
recon: block(&[(0, 74), (1, 44), (2, -44)]),
last: 5,
},
Case {
q: 64,
coeffs: block(&[(0, 130), (1, 41), (2, -39), (8, 8)]),
first: 0,
ctx0: 0,
levels: block(&[(0, 2)]),
recon: block(&[(0, 118)]),
last: 0,
},
Case {
q: 100,
coeffs: block(&[(1, 3), (4, -2), (5, 4)]),
first: 1,
ctx0: 0,
levels: [0; 16],
recon: [0; 16],
last: 0,
},
Case {
q: 52,
coeffs: block(&[(0, 200), (1, 55), (3, 48), (7, 40), (12, 35)]),
first: 0,
ctx0: 0,
levels: block(&[(0, 4), (1, 1)]),
recon: block(&[(0, 188), (1, 56)]),
last: 1,
},
];
for (i, c) in cases.iter().enumerate() {
let out = trellis(c.q, c.coeffs, c.first, c.ctx0);
assert_eq!(out.levels, c.levels, "case {i} levels");
assert_eq!(out.recon, c.recon, "case {i} recon");
assert_eq!(out.last, c.last, "case {i} last");
}
}
#[test]
fn trellis_tie_break_keeps_the_earlier_termination() {
let coeffs = block(&[(0, 33), (3, 20)]);
let out = trellis(17, coeffs, 0, 0);
assert_eq!(
out.levels,
block(&[(0, 2)]),
"only the DC is coded (level 2)"
);
assert_eq!(out.recon, block(&[(0, 38)]), "recon is DC-only (2 * 19)");
assert_eq!(
out.last, 0,
"the strict-< tie-break keeps the earlier (DC-only) termination"
);
}
fn brute_force(
coeffs: [i16; 16],
pair: super::QPair,
first: usize,
ctx0: usize,
lambda: i64,
) -> super::Quantized {
let npos = 16 - first;
let mut cand = [[0i32; 3]; 16];
let mut ncand = [0usize; 16];
for k in 0..npos {
let n = first + k;
let j = ZIGZAG[n];
let factor = if j == 0 { pair.dc } else { pair.ac };
let l0 = factor.quantize(i32::from(coeffs[j])).abs();
let mut list = [l0, 0, 0];
let mut c = 1usize;
if l0 >= 1 {
list[c] = l0 - 1;
c += 1;
}
if l0 >= 2 {
list[c] = 0;
c += 1;
}
cand[k] = list;
ncand[k] = c;
}
let mut idx = [0usize; 16];
let mut best_total = i64::MAX;
let mut best = super::Quantized {
levels: [0; 16],
recon: [0; 16],
last: first as i32 - 1,
};
loop {
let mut levels = [0i16; 16];
let mut recon = [0i16; 16];
let mut last = first as i32 - 1;
let mut dist = 0i64;
for k in 0..npos {
let n = first + k;
let j = ZIGZAG[n];
let factor = if j == 0 { pair.dc } else { pair.ac };
let q = factor.q;
let mag = cand[k][idx[k]];
let abs_c = i32::from(coeffs[j]).abs();
let err = i64::from(abs_c - mag * q);
dist += err * err;
let signed = if coeffs[j] < 0 { -mag } else { mag };
levels[n] = signed as i16;
recon[j] = (signed * q) as i16;
if mag != 0 {
last = n as i32;
}
}
let rate = block_token_cost(levels, first, last, 0, ctx0, &COEFFS_PROBA_0);
let total = super::RD_DISTO_MULT * dist + lambda * rate;
if total < best_total {
best_total = total;
best = super::Quantized {
levels,
recon,
last,
};
}
let mut k = 0;
loop {
if k == npos {
return best;
}
idx[k] += 1;
if idx[k] < ncand[k] {
break;
}
idx[k] = 0;
k += 1;
}
}
}
#[test]
fn trellis_matches_the_brute_force_optimum() {
let mut s: u32 = 0x1234_5678;
let mut next = || {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
s
};
for trial in 0..400u32 {
let base_q = 8 + (next() % 100) as i32;
let first = (next() % 2) as usize; let ctx0 = (next() % 3) as usize;
let span = first + 5 + (next() % 3) as usize; let mut coeffs = [0i16; 16];
for n in first..16 {
let active = n < span || (n == 15 && (next() % 4 == 0));
if active {
let mag = (next() % 260) as i32 - 130;
coeffs[ZIGZAG[n]] = mag as i16;
}
}
let q = Quantizer::new(base_q);
let lambda = trellis_lambda(q.y1.ac.q);
let got = trellis_quantize_block(coeffs, q.y1, first, ctx0, 0, &COEFFS_PROBA_0, lambda);
let want = brute_force(coeffs, q.y1, first, ctx0, lambda);
assert_eq!(
got.last, want.last,
"trial {trial}: last mismatch (q={base_q}, first={first}, ctx0={ctx0})"
);
assert_eq!(got.levels, want.levels, "trial {trial}: levels mismatch");
assert_eq!(got.recon, want.recon, "trial {trial}: recon mismatch");
}
}
}