#[test]
fn test_celt_mdct_passthrough() {
use rusty_opus::modes::default_mode;
let mode = default_mode();
let frame_size = 960;
let overlap = mode.overlap;
let shift = 0usize;
let b = 1usize;
let mdct_input_size = mode.mdct.n + overlap;
let syn_mem_size = mdct_input_size + frame_size;
let freq_hz = 440.0 / 48000.0 * 2.0 * std::f32::consts::PI;
let total_samples = 5 * frame_size;
let mut all_in = vec![0.0f32; total_samples];
for i in 0..total_samples {
all_in[i] = (freq_hz * i as f32).sin();
}
let mut syn_mem = vec![0.0f32; syn_mem_size];
let decode_buffer_size = mode.mdct.n + overlap;
let mut decode_mem = vec![0.0f32; decode_buffer_size];
let mut all_out = Vec::new();
for frame_idx in 0..3 {
let frame_start = (frame_idx + 1) * frame_size;
for i in 0..syn_mem_size - frame_size {
syn_mem[i] = syn_mem[i + frame_size];
}
for i in 0..frame_size {
syn_mem[syn_mem_size - frame_size + i] = all_in[frame_start + i];
}
let mut freq_buf = vec![0.0f32; frame_size];
mode.mdct.forward(
&syn_mem[syn_mem_size - mdct_input_size..],
&mut freq_buf,
mode.window,
overlap,
shift,
b,
);
mode.mdct
.backward(&freq_buf, &mut decode_mem, mode.window, overlap, shift, b);
let mut frame_out = vec![0.0f32; frame_size];
frame_out.copy_from_slice(&decode_mem[overlap..overlap + frame_size]);
all_out.extend_from_slice(&frame_out);
}
let compare_frame = 2;
let mut best_snr = -100.0f64;
let mut best_delay = 0usize;
for delay in 0..2 * frame_size {
let mut signal_power = 0.0f64;
let mut error_power = 0.0f64;
let mut count = 0;
for i in 0..frame_size {
let out_idx = compare_frame * frame_size + i;
if delay > (compare_frame + 1) * frame_size + i {
continue;
}
let in_idx = (compare_frame + 1) * frame_size + i - delay;
if in_idx < total_samples && out_idx < all_out.len() {
let sig = all_in[in_idx] as f64;
let out = all_out[out_idx] as f64;
signal_power += sig * sig;
error_power += (out - sig) * (out - sig);
count += 1;
}
}
if count > frame_size / 2 && signal_power > 1e-10 {
let snr = 10.0 * (signal_power / error_power.max(1e-20)).log10();
if snr > best_snr {
best_snr = snr;
best_delay = delay;
}
}
}
eprintln!(
"CELT MDCT passthrough: best SNR = {:.2} dB at delay = {}",
best_snr, best_delay
);
assert!(
best_snr > 0.0,
"MDCT passthrough SNR too low: {:.2} dB at delay {}",
best_snr,
best_delay
);
}