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)
})
}
fn lower_magnitudes(
coef_natural: &[i16; 64],
q_block: &mut [i16; 64],
table: &QuantizationTable,
ac_table: &HuffmanTable,
lambda: f64,
) -> u32 {
let mut changed = 0;
let mut prev = 0usize;
for i in 1..64 {
let q = q_block[i];
if q == 0 {
continue;
}
let run_before = (i - prev - 1) as u32;
prev = i;
if q.abs() < 2 {
continue;
}
let lowered = q - q.signum();
let z = ZIGZAG[i] as usize & 0x3f;
let step = table.divisor(z) as f64;
let c = coef_natural[z] as f64;
let err_now = c - q as f64 * step;
let err_low = c - lowered as f64 * step;
let delta_d = err_low * err_low - err_now * err_now;
let r = (run_before % 16) as u8;
let (size_now, _) = get_code(q);
let (size_low, _) = get_code(lowered);
let rate_now = ac_table.code_len((r << 4) | size_now) as f64 + size_now as f64;
let rate_low = ac_table.code_len((r << 4) | size_low) as f64 + size_low as f64;
if delta_d + lambda * (rate_low - rate_now) < 0.0 {
q_block[i] = lowered;
changed += 1;
}
}
changed
}
#[allow(clippy::needless_range_loop)]
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;
}
}
let mut zeroed = 0;
if best_k < nnz {
for k in best_k..nnz {
q_block[pos[k] as usize] = 0;
}
zeroed = (nnz - best_k) as u32;
}
if magnitudes_enabled() {
zeroed += lower_magnitudes(coef_natural, q_block, table, ac_table, lambda);
}
zeroed
}
fn magnitudes_enabled() -> bool {
use std::sync::OnceLock;
static V: OnceLock<bool> = OnceLock::new();
*V.get_or_init(|| {
std::env::var("RUSTY_JPEG_TRELLIS_MAG")
.map(|v| v != "0")
.unwrap_or(true)
})
}