use crate::av1_coder::*;
use crate::coeffs::get_lo_ctx_2d;
use crate::cost::*;
use crate::tables::{COEFF_BASE_RANGE, LO_CTX_OFF, NUM_BASE_LEVELS, level_byte};
#[inline]
fn trellis_lambda_scale() -> f64 {
2.0
}
#[allow(clippy::too_many_arguments, clippy::type_complexity)]
pub(crate) fn trellis_optimize_ctx(
cf: &mut [i32],
tf: &[f64],
dc_q: f64,
ac_q: f64,
scan: &[usize],
lambda0: f64,
w: usize,
cdfs: &Cdfs,
cls: usize,
plane: usize,
eob_bin_cdf: &[u16],
dcs_ctx: usize,
) {
if lambda0 <= 0.0 {
return;
}
let n = scan.len();
let lambda = lambda0 * ac_q * ac_q * trellis_lambda_scale();
let log2w = w.trailing_zeros() as usize;
let stride = w;
let base_tok = &cdfs.base_tok[cls][plane];
let br_tok = &cdfs.br_tok[cls][plane];
let eob_hi = &cdfs.eob_hi[cls][plane];
let eob_base = &cdfs.eob_base[cls][plane];
let dc_sign = &cdfs.dc_sign[plane];
let dq2_dc = dc_q * dc_q;
let dq2_ac = ac_q * ac_q;
let dist = |rc: usize, lev: i32| {
let dq2 = if rc == 0 { dq2_dc } else { dq2_ac };
let e = tf[rc].abs() - (lev.abs() as f64);
dq2 * (e * e)
};
let use_br_table = n >= 256;
let mut br_cum = [[0f32; 13]; 21];
if use_br_table {
for (row, br) in br_cum.iter_mut().zip(br_tok.iter()) {
let c = [
cdf_cost(br, 0),
cdf_cost(br, 1),
cdf_cost(br, 2),
cdf_cost(br, 3),
];
for (j, slot) in row.iter_mut().enumerate() {
let mut coded = 0i32;
let mut bits = 0.0f32;
for _ in 0..(COEFF_BASE_RANGE / 3) {
let s = (j as i32 - coded).min(3);
bits += c[s as usize];
coded += s;
if s < 3 {
break;
}
}
*slot = bits;
}
}
}
let hi_cost = |m: u32, bc: usize| -> f32 {
if use_br_table {
let total_br = (m as i32 - (NUM_BASE_LEVELS + 1)).min(COEFF_BASE_RANGE);
let mut bits = br_cum[bc][total_br as usize];
if m >= 15 {
bits += golomb_cost(m - 15);
}
bits
} else {
hi_tok_cost(m, &br_tok[bc])
}
};
let eob: i32 = scan
.iter()
.rposition(|&rc| cf[rc] != 0)
.map_or(-1, |i| i as i32);
if eob < 0 {
return;
}
let eu = eob as usize;
thread_local! {
static SCRATCH: std::cell::RefCell<(Vec<u8>, Vec<f64>, Vec<f64>, Vec<f32>)> =
const { std::cell::RefCell::new((Vec::new(), Vec::new(), Vec::new(), Vec::new())) };
}
let (mut levels, mut pre, mut suf0, mut irate) = SCRATCH.with(|s| {
let mut b = s.borrow_mut();
(
std::mem::take(&mut b.0),
std::mem::take(&mut b.1),
std::mem::take(&mut b.2),
std::mem::take(&mut b.3),
)
});
levels.clear();
levels.resize(w * (w + 4), 0);
let set_level = |levels: &mut [u8], rc: usize, m: u32| {
levels[(rc >> log2w) * stride + (rc & (w - 1))] = level_byte(m);
};
for &rc in &scan[..eu + 1] {
set_level(&mut levels, rc, cf[rc].unsigned_abs());
}
let interior_ctx = |levels: &[u8], rc: usize| -> (usize, usize) {
let (x, y) = (rc >> log2w, rc & (w - 1));
let (ctx, hi_mag) = get_lo_ctx_2d(levels, x, y, &LO_CTX_OFF, stride);
let mag = hi_mag & 63;
let bc = (if (y | x) > 1 { 14 } else { 7 }) + if mag > 12 { 6 } else { (mag + 1) >> 1 };
(ctx, bc as usize)
};
let dc_brc = |levels: &[u8]| -> usize {
let mag = (levels[1] as u32 + levels[stride] as u32 + levels[stride + 1] as u32) & 63;
if mag > 12 {
6
} else {
((mag + 1) >> 1) as usize
}
};
let interior_rate = |ctx: usize, bc: usize, k: u32| -> f32 {
if k == 0 {
return cdf_cost(&base_tok[ctx], 0);
}
let tok = k.min(3);
let mut b = cdf_cost(&base_tok[ctx], tok as usize);
if tok == 3 {
b += hi_cost(k, bc);
}
b + 1.0 };
for i in (1..(eob as usize)).rev() {
let rc = scan[i];
let l = cf[rc].unsigned_abs();
if l == 0 {
continue;
}
let (ctx, bc) = interior_ctx(&levels, rc);
let bt = &base_tok[ctx];
let bt0 = cdf_cost(bt, 0);
let bt1 = cdf_cost(bt, 1);
let bt2 = cdf_cost(bt, 2);
let bt3 = cdf_cost(bt, 3);
let rate_k = |k: u32| -> f32 {
match k {
0 => bt0,
1 => bt1 + 1.0,
2 => bt2 + 1.0,
_ => (bt3 + hi_cost(k, bc)) + 1.0,
}
};
let mut best_k = l;
let mut best_c = dist(rc, l as i32) + rate_cost(lambda, rate_k(l));
for k in (0..l).rev() {
let dk = dist(rc, k as i32);
if dk >= best_c {
break;
}
let c = dk + rate_cost(lambda, rate_k(k));
if c < best_c {
best_c = c;
best_k = k;
}
}
if best_k != l {
cf[rc] = if cf[rc] < 0 {
-(best_k as i32)
} else {
best_k as i32
};
set_level(&mut levels, rc, best_k);
}
}
{
let rc = scan[0];
let l = cf[rc].unsigned_abs();
if l != 0 {
let bc = dc_brc(&levels);
let sgn = (cf[rc] < 0) as usize;
let dc_rate = |k: u32| -> f32 {
if k == 0 {
return cdf_cost(&base_tok[0], 0);
}
let tok = k.min(3);
let mut b = cdf_cost(&base_tok[0], tok as usize);
if tok == 3 {
b += hi_cost(k, bc);
}
b + cdf_cost(&dc_sign[dcs_ctx], sgn)
};
let mut best_k = l;
let mut best_c = dist(rc, l as i32) + rate_cost(lambda, dc_rate(l));
for k in (0..l).rev() {
let dk = dist(rc, k as i32);
if dk >= best_c {
break;
}
let c = dk + rate_cost(lambda, dc_rate(k));
if c < best_c {
best_c = c;
best_k = k;
}
}
if best_k != l {
cf[rc] = if cf[rc] < 0 {
-(best_k as i32)
} else {
best_k as i32
};
set_level(&mut levels, rc, best_k);
}
}
}
let eob_pt_cost = |e: usize| -> f32 {
let bin = if e < 2 {
e
} else {
32 - (e as u32).leading_zeros() as usize
};
let mut c = cdf_cost(eob_bin_cdf, bin);
if bin > 1 {
let nbits = bin - 2;
c += cdf_cost(&eob_hi[bin], (e >> nbits) & 1);
c += nbits as f32; }
c
};
let eob_coeff_cost = |e: usize, m: u32| -> f32 {
let ctx_e = 1 + (e > n / 8) as usize + (e > n / 4) as usize;
let tok = m.min(3);
let mut c = cdf_cost(&eob_base[ctx_e], tok as usize - 1);
if tok == 3 {
let rc = scan[e];
let (ex, ey) = (rc >> log2w, rc & (w - 1));
let bc = if (ex | ey) > 1 { 14 } else { 7 };
c += hi_cost(m, bc);
}
c + 1.0 };
pre.resize(n + 1, 0.0);
irate.resize(n, 0.0);
let mut acc = 0.0f64; for ((&rc, ir), p) in scan[1..eu + 1]
.iter()
.zip(irate[1..eu + 1].iter_mut())
.zip(pre[2..eu + 2].iter_mut())
{
let (ctx, bc) = interior_ctx(&levels, rc);
let r = interior_rate(ctx, bc, cf[rc].unsigned_abs());
*ir = r;
acc = (acc + rate_cost(lambda, r)) + dist(rc, cf[rc]);
*p = acc;
}
for (&rc, p) in scan[eu + 1..n].iter().zip(pre[eu + 2..n + 1].iter_mut()) {
acc += dist(rc, 0);
*p = acc;
}
suf0.resize(n + 1, 0.0);
suf0[n] = 0.0; let mut sacc = 0.0f64;
for (&rc, s) in scan[1..n].iter().rev().zip(suf0[1..n].iter_mut().rev()) {
sacc += dist(rc, 0);
*s = sacc;
}
let dc_rc = scan[0];
let dc_m = cf[dc_rc].unsigned_abs();
let dc_cost = if dc_m == 0 {
rate_cost(lambda, cdf_cost(&base_tok[0], 0))
} else {
let bc = dc_brc(&levels);
let tok = dc_m.min(3);
let mut b = cdf_cost(&base_tok[0], tok as usize);
if tok == 3 {
b += hi_cost(dc_m, bc);
}
b += cdf_cost(&dc_sign[dcs_ctx], (cf[dc_rc] < 0) as usize);
rate_cost(lambda, b)
} + dist(dc_rc, cf[dc_rc]);
let mut best_e: i32 = -1;
let mut best_cost = f64::INFINITY;
assert!(scan.len() >= n, "scan must be indexed up to n-1");
assert!(irate.len() >= n, "irate must be indexed up to n");
assert!(pre.len() > n, "pre must be indexed up to n+1");
assert!(suf0.len() > n, "suf0 must be indexed up to n+1");
for e in 1..n {
let rc = scan[e];
if cf[rc] == 0 {
continue; }
let interior_e = rate_cost(lambda, irate[e]);
let c = dc_cost
+ (pre[e + 1] - interior_e)
+ rate_cost(
lambda,
eob_pt_cost(e) + eob_coeff_cost(e, cf[rc].unsigned_abs()),
)
+ suf0[e + 1];
if c < best_cost {
best_cost = c;
best_e = e as i32;
}
}
if dc_m != 0 {
let ctx_e = 1usize; let tok = dc_m.min(3);
let mut c0 = cdf_cost(eob_bin_cdf, 0) + cdf_cost(&eob_base[ctx_e], tok as usize - 1);
if tok == 3 {
c0 += hi_cost(dc_m, dc_brc(&levels));
}
c0 += cdf_cost(&dc_sign[dcs_ctx], (cf[dc_rc] < 0) as usize);
let total0 = rate_cost(lambda, c0) + dist(dc_rc, cf[dc_rc]) + suf0[1];
if total0 < best_cost {
best_cost = total0;
best_e = 0;
}
}
let skip_cost = suf0[1] + dist(dc_rc, 0) + rate_cost(lambda, 1.0f32);
if best_e < 0 || skip_cost < best_cost {
for &rc in scan.iter() {
cf[rc] = 0;
}
} else {
for i in (best_e as usize + 1)..n {
cf[scan[i]] = 0;
}
}
SCRATCH.with(|s| {
let mut b = s.borrow_mut();
b.0 = std::mem::take(&mut levels);
b.1 = std::mem::take(&mut pre);
b.2 = std::mem::take(&mut suf0);
b.3 = std::mem::take(&mut irate);
});
}
pub(crate) fn trellis_optimize(
cf: &mut [i32],
tf: &[f64],
dc_q: f64,
ac_q: f64,
scan: &[usize],
lambda0: f64,
) {
if lambda0 <= 0.0 {
return; }
let n = scan.len();
let lambda = lambda0 * ac_q * ac_q * trellis_lambda_scale();
let (dc_q2, ac_q2) = (dc_q * dc_q, ac_q * ac_q);
let d = |rc: usize, lev: i32| {
let dq2 = if rc == 0 { dc_q2 } else { ac_q2 };
let e = tf[rc].abs() - lev.unsigned_abs() as f64;
dq2 * e * e
};
let mut eob_idx: i32 = -1;
for (i, &x) in scan[..n].iter().enumerate() {
if cf[x] != 0 {
eob_idx = i as i32;
}
}
if eob_idx < 0 {
return; }
for &rc in scan[..=eob_idx as usize].iter() {
let c = cf[rc];
if c == 0 {
continue;
}
let l = c.unsigned_abs();
let dq2 = if rc == 0 { dc_q2 } else { ac_q2 };
let at = tf[rc].abs();
let (e_l, e_dn) = (at - l as f64, at - (l - 1) as f64);
let cost_l = dq2 * e_l * e_l + rate_cost(lambda, coef_rate_bits(l));
let cost_dn = dq2 * e_dn * e_dn + rate_cost(lambda, coef_rate_bits(l - 1));
if cost_dn < cost_l {
cf[rc] = if c < 0 { -(l as i32 - 1) } else { l as i32 - 1 };
}
}
thread_local! {
static SCRATCH: std::cell::RefCell<(Vec<f64>, Vec<f64>)> =
const { std::cell::RefCell::new((Vec::new(), Vec::new())) };
}
let (mut suf0, mut pre) = SCRATCH.with(|s| {
let mut b = s.borrow_mut();
(std::mem::take(&mut b.0), std::mem::take(&mut b.1))
});
suf0.resize(n + 1, 0.0); suf0[n] = 0.0; assert!(suf0.len() > n, "suf0 must be indexed up to n");
for (i, &s) in (0..n).rev().zip(scan[..n].iter().rev()) {
suf0[i] = suf0[i + 1] + d(s, 0);
}
pre.resize(n + 1, 0.0); assert!(pre.len() > n, "pre must be indexed up to n");
pre[0] = 0.0; for (i, &rc) in scan[..n].iter().enumerate() {
pre[i + 1] =
pre[i] + d(rc, cf[rc]) + rate_cost(lambda, coef_rate_bits(cf[rc].unsigned_abs()));
}
let eob_sig = |e: usize| -> f32 {
let bin = if e < 2 {
e
} else {
(32 - (e as u32).leading_zeros()) as usize
};
let extra = if bin > 1 { bin - 2 } else { 0 };
(bin as f32) * 0.9 + extra as f32 + 2.0 };
let mut best_e: i32 = -1;
let mut best_cost = f64::INFINITY;
for (e, ((&rc, &pre), &suf0)) in scan[..n]
.iter()
.zip(pre.iter())
.zip(suf0[1..].iter())
.enumerate()
{
if cf[rc] == 0 {
continue; }
let c = pre + d(rc, cf[rc]) + rate_cost(lambda, eob_sig(e)) + suf0;
if c < best_cost {
best_cost = c;
best_e = e as i32;
}
}
let skip_cost = suf0[0] + rate_cost(lambda, 1.0f32); if best_e < 0 || skip_cost < best_cost {
for &rc in scan.iter() {
cf[rc] = 0;
}
} else {
for &x in scan[(best_e as usize + 1)..n].iter() {
cf[x] = 0;
}
}
SCRATCH.with(|s| {
let mut b = s.borrow_mut();
b.0 = std::mem::take(&mut suf0);
b.1 = std::mem::take(&mut pre);
});
}