use crate::encode::huffman::HuffmanTable;
use crate::encode::quantization::QuantizationTable;
use crate::encode::writer::{get_code, ZIGZAG};
const DEFAULT_LAMBDA_SCALE: f32 = 0.10;
pub(crate) fn lambda_scale() -> f32 {
use std::sync::OnceLock;
static V: OnceLock<f32> = OnceLock::new();
*V.get_or_init(|| {
std::env::var("RUSTY_JPEG_TRELLIS_LAMBDA")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(DEFAULT_LAMBDA_SCALE)
})
}
pub(crate) fn truncate_rd(
coef_natural: &[i16; 64],
q_block: &mut [i16; 64],
table: &QuantizationTable,
ac_table: &HuffmanTable,
) -> u32 {
let mut pos = [0u8; 63];
let mut nnz = 0usize;
for i in 1..64 {
if q_block[i] != 0 {
pos[nnz] = i as u8;
nnz += 1;
}
}
if nnz == 0 {
return 0;
}
let bits = |sym: u8| -> u32 { ac_table.code_len(sym) as u32 };
let eob_bits = bits(0x00);
let mut prefix = [0u32; 64];
let mut acc = 0u32;
let mut prev = 0usize;
for k in 0..nnz {
let p = pos[k] as usize;
let mut r = (p - prev - 1) as u32;
while r > 15 {
acc += bits(0xF0);
r -= 16;
}
let (size, _) = get_code(q_block[p]);
acc += bits(((r as u8) << 4) | size) + size as u32;
prefix[k + 1] = acc;
prev = p;
}
let mut drop_cost = [0.0f64; 64];
for k in (0..nnz).rev() {
let p = pos[k] as usize;
let z = ZIGZAG[p] as usize & 0x3f;
let step = table.divisor(z) as f64;
let c = coef_natural[z] as f64;
let kept_err = c - q_block[p] as f64 * step;
let d = c * c - kept_err * kept_err;
drop_cost[k] = drop_cost[k + 1] + d.max(0.0);
}
let mut mean_sq = 0.0f64;
for k in 0..nnz {
let z = ZIGZAG[pos[k] as usize] as usize & 0x3f;
let s = table.divisor(z) as f64;
mean_sq += s * s;
}
mean_sq /= nnz as f64;
let lambda = lambda_scale() as f64 * mean_sq;
let mut best_k = nnz;
let mut best = f64::INFINITY;
for k in 0..=nnz {
let last_pos = if k == 0 { 0 } else { pos[k - 1] as usize };
let rate = prefix[k] + if last_pos < 63 { eob_bits } else { 0 };
let cost = drop_cost[k] + lambda * rate as f64;
if cost < best {
best = cost;
best_k = k;
}
}
if best_k == nnz {
return 0;
}
for k in best_k..nnz {
q_block[pos[k] as usize] = 0;
}
(nnz - best_k) as u32
}