use crate::EvaluationError;
use scirs2_core::ndarray::{Array1, Array2};
use scirs2_fft::{RealFftPlanner, RealToComplex};
use std::f32::consts::PI;
use std::sync::Mutex;
use voirs_sdk::AudioBuffer;
pub struct PESQEvaluator {
sample_rate: u32,
fft_planner: Mutex<RealFftPlanner<f32>>,
bark_mapping: Vec<f32>,
perceptual_weights: Array1<f32>,
}
impl PESQEvaluator {
pub fn new_narrowband() -> Result<Self, EvaluationError> {
Self::new(8000)
}
pub fn new_wideband() -> Result<Self, EvaluationError> {
Self::new(16000)
}
fn new(sample_rate: u32) -> Result<Self, EvaluationError> {
if sample_rate != 8000 && sample_rate != 16000 {
return Err(EvaluationError::InvalidInput {
message: "PESQ only supports 8 kHz and 16 kHz sample rates".to_string(),
});
}
let fft_planner = Mutex::new(RealFftPlanner::<f32>::new());
let bark_mapping = Self::create_bark_mapping(sample_rate);
let perceptual_weights = Self::create_perceptual_weights(&bark_mapping);
Ok(Self {
sample_rate,
fft_planner,
bark_mapping,
perceptual_weights,
})
}
pub async fn calculate_pesq(
&self,
reference: &AudioBuffer,
degraded: &AudioBuffer,
) -> Result<f32, EvaluationError> {
self.validate_inputs(reference, degraded)?;
let (ref_aligned, deg_aligned) = self.level_alignment(reference, degraded)?;
let ref_filtered = self.input_filtering(&ref_aligned)?;
let deg_filtered = self.input_filtering(°_aligned)?;
let (ref_time_aligned, deg_time_aligned) =
self.time_alignment(&ref_filtered, °_filtered)?;
let ref_bark = self.auditory_transform(&ref_time_aligned)?;
let deg_bark = self.auditory_transform(°_time_aligned)?;
let disturbance_frames = self.cognitive_modeling(&ref_bark, °_bark)?;
let pesq_score = self.calculate_final_score(&disturbance_frames)?;
Ok(pesq_score)
}
fn validate_inputs(
&self,
reference: &AudioBuffer,
degraded: &AudioBuffer,
) -> Result<(), EvaluationError> {
if reference.sample_rate() != self.sample_rate {
return Err(EvaluationError::InvalidInput {
message: crate::error_enhancement::sample_rate_mismatch_error(
"PESQ",
self.sample_rate,
reference.sample_rate(),
"reference",
),
});
}
if degraded.sample_rate() != self.sample_rate {
return Err(EvaluationError::InvalidInput {
message: crate::error_enhancement::sample_rate_mismatch_error(
"PESQ",
self.sample_rate,
degraded.sample_rate(),
"degraded",
),
});
}
if reference.channels() != 1 {
return Err(EvaluationError::InvalidInput {
message: crate::error_enhancement::channel_mismatch_error(
"PESQ",
1,
reference.channels(),
"reference",
),
});
}
if degraded.channels() != 1 {
return Err(EvaluationError::InvalidInput {
message: crate::error_enhancement::channel_mismatch_error(
"PESQ",
1,
degraded.channels(),
"degraded",
),
});
}
Ok(())
}
fn level_alignment(
&self,
reference: &AudioBuffer,
degraded: &AudioBuffer,
) -> Result<(Array1<f32>, Array1<f32>), EvaluationError> {
let ref_samples = Array1::from_vec(reference.samples().to_vec());
let deg_samples = Array1::from_vec(degraded.samples().to_vec());
let ref_mean_sq = ref_samples.mapv(|x| x * x).mean().unwrap_or(0.0);
let deg_mean_sq = deg_samples.mapv(|x| x * x).mean().unwrap_or(0.0);
let epsilon = 1e-12f32;
let ref_rms = (ref_mean_sq.max(0.0) + epsilon).sqrt();
let deg_rms = (deg_mean_sq.max(0.0) + epsilon).sqrt();
if ref_rms <= epsilon || deg_rms <= epsilon {
return Err(EvaluationError::AudioProcessingError {
message: "Audio contains no signal".to_string(),
source: None,
});
}
let target_level = 10_f32.powf(-26.0 / 20.0);
let ref_scale = target_level / ref_rms;
let deg_scale = target_level / deg_rms;
let ref_aligned = ref_samples.mapv(|x| x * ref_scale);
let deg_aligned = deg_samples.mapv(|x| x * deg_scale);
Ok((ref_aligned, deg_aligned))
}
fn input_filtering(&self, signal: &Array1<f32>) -> Result<Array1<f32>, EvaluationError> {
let irs_filtered = self.apply_irs_filter(signal)?;
let filtered = self.apply_receiving_filter(&irs_filtered)?;
Ok(filtered)
}
fn apply_irs_filter(&self, signal: &Array1<f32>) -> Result<Array1<f32>, EvaluationError> {
let (b_coeffs, a_coeffs) = if self.sample_rate == 8000 {
(
vec![0.008_378_7, 0.025_136_1, 0.025_136_1, 0.008_378_7],
vec![1.0, -1.760_041_8, 0.890_470_1, -0.160_612_1],
)
} else {
(
vec![0.001_687_87, 0.005_073_61, 0.005_073_61, 0.001_687_87],
vec![1.0, -2.760_041_8, 2.590_47, -0.860_612_1],
)
};
self.apply_biquad_filter(signal, &b_coeffs, &a_coeffs)
}
fn apply_receiving_filter(&self, signal: &Array1<f32>) -> Result<Array1<f32>, EvaluationError> {
let (b_coeffs, a_coeffs) = if self.sample_rate == 8000 {
(vec![1.0, -2.0, 1.0], vec![1.0, -1.9878, 0.9881])
} else {
(vec![1.0, -2.0, 1.0], vec![1.0, -1.9939, 0.9940])
};
self.apply_biquad_filter(signal, &b_coeffs, &a_coeffs)
}
fn apply_biquad_filter(
&self,
signal: &Array1<f32>,
b_coeffs: &[f32],
a_coeffs: &[f32],
) -> Result<Array1<f32>, EvaluationError> {
let len = signal.len();
let mut filtered = Array1::zeros(len);
let mut x_delay = vec![0.0; b_coeffs.len()];
let mut y_delay = vec![0.0; a_coeffs.len()];
for i in 0..len {
for j in (1..x_delay.len()).rev() {
x_delay[j] = x_delay[j - 1];
}
for j in (1..y_delay.len()).rev() {
y_delay[j] = y_delay[j - 1];
}
x_delay[0] = signal[i];
let mut output = 0.0;
for j in 0..b_coeffs.len() {
output += b_coeffs[j] * x_delay[j];
}
for j in 1..a_coeffs.len() {
output -= a_coeffs[j] * y_delay[j];
}
y_delay[0] = output;
filtered[i] = output;
}
Ok(filtered)
}
fn time_alignment(
&self,
reference: &Array1<f32>,
degraded: &Array1<f32>,
) -> Result<(Array1<f32>, Array1<f32>), EvaluationError> {
let max_delay = (0.5 * self.sample_rate as f32) as usize; let min_length = 8 * self.sample_rate as usize;
let delay = self.find_optimal_delay(reference, degraded, max_delay)?;
let (ref_aligned, deg_aligned) = if delay >= 0 {
let delay = delay as usize;
let start_ref = 0;
let start_deg = delay;
if start_deg >= degraded.len() {
return Ok((reference.clone(), degraded.clone()));
}
let max_length = (reference.len() - start_ref).min(degraded.len() - start_deg);
let length = max_length.max(min_length.min(max_length));
let ref_end = (start_ref + length).min(reference.len());
let deg_end = (start_deg + length).min(degraded.len());
let actual_length = (ref_end - start_ref).min(deg_end - start_deg);
let ref_slice = reference
.slice(scirs2_core::ndarray::s![
start_ref..start_ref + actual_length
])
.to_owned();
let deg_slice = degraded
.slice(scirs2_core::ndarray::s![
start_deg..start_deg + actual_length
])
.to_owned();
(ref_slice, deg_slice)
} else {
let delay = (-delay) as usize;
let start_ref = delay;
let start_deg = 0;
if start_ref >= reference.len() {
return Ok((reference.clone(), degraded.clone()));
}
let max_length = (reference.len() - start_ref).min(degraded.len() - start_deg);
let length = max_length.max(min_length.min(max_length));
let ref_end = (start_ref + length).min(reference.len());
let deg_end = (start_deg + length).min(degraded.len());
let actual_length = (ref_end - start_ref).min(deg_end - start_deg);
let ref_slice = reference
.slice(scirs2_core::ndarray::s![
start_ref..start_ref + actual_length
])
.to_owned();
let deg_slice = degraded
.slice(scirs2_core::ndarray::s![
start_deg..start_deg + actual_length
])
.to_owned();
(ref_slice, deg_slice)
};
Ok((ref_aligned, deg_aligned))
}
fn find_optimal_delay(
&self,
reference: &Array1<f32>,
degraded: &Array1<f32>,
max_delay: usize,
) -> Result<i32, EvaluationError> {
let window_size = 4 * self.sample_rate as usize; let ref_len = reference.len().min(window_size);
let deg_len = degraded.len().min(window_size);
let mut best_correlation = f32::NEG_INFINITY;
let mut best_delay = 0i32;
for delay in -(max_delay as i32)..=(max_delay as i32) {
let correlation = if delay >= 0 {
let delay = delay as usize;
if delay >= deg_len {
continue;
}
let length = (ref_len).min(deg_len - delay);
if length < self.sample_rate as usize {
continue;
}
let ref_segment = reference.slice(scirs2_core::ndarray::s![0..length]);
let deg_segment = degraded.slice(scirs2_core::ndarray::s![delay..delay + length]);
self.calculate_correlation(&ref_segment, °_segment)
} else {
let delay = (-delay) as usize;
if delay >= ref_len {
continue;
}
let length = (ref_len - delay).min(deg_len);
if length < self.sample_rate as usize {
continue;
}
let ref_segment = reference.slice(scirs2_core::ndarray::s![delay..delay + length]);
let deg_segment = degraded.slice(scirs2_core::ndarray::s![0..length]);
self.calculate_correlation(&ref_segment, °_segment)
};
if correlation > best_correlation {
best_correlation = correlation;
best_delay = delay;
}
}
Ok(best_delay)
}
fn calculate_correlation(
&self,
signal1: &scirs2_core::ndarray::ArrayView1<f32>,
signal2: &scirs2_core::ndarray::ArrayView1<f32>,
) -> f32 {
let mean1 = signal1.mean().unwrap_or(0.0);
let mean2 = signal2.mean().unwrap_or(0.0);
let mut numerator = 0.0;
let mut sum_sq1 = 0.0;
let mut sum_sq2 = 0.0;
for (&s1, &s2) in signal1.iter().zip(signal2.iter()) {
let diff1 = s1 - mean1;
let diff2 = s2 - mean2;
numerator += diff1 * diff2;
sum_sq1 += diff1 * diff1;
sum_sq2 += diff2 * diff2;
}
let denominator = (sum_sq1 * sum_sq2).sqrt();
if denominator > 1e-12f32 {
(numerator / denominator).clamp(-1.0, 1.0) } else {
0.0
}
}
fn auditory_transform(&self, signal: &Array1<f32>) -> Result<Array2<f32>, EvaluationError> {
let frame_size = 512;
let hop_size = 256;
let num_frames = (signal.len() - frame_size) / hop_size + 1;
let num_bark_bands = self.bark_mapping.len();
let mut bark_spectrum = Array2::zeros((num_frames, num_bark_bands));
let mut fft_planner = self
.fft_planner
.lock()
.expect("lock should not be poisoned");
let fft = fft_planner.plan_fft_forward(frame_size);
let mut spectrum = vec![scirs2_core::Complex::new(0.0, 0.0); fft.output_len()];
for (frame_idx, frame_start) in (0..signal.len() - frame_size + 1)
.step_by(hop_size)
.enumerate()
{
if frame_idx >= num_frames {
break;
}
let mut frame = Array1::zeros(frame_size);
for i in 0..frame_size {
if frame_start + i < signal.len() {
let window =
0.5 * (1.0 - (2.0 * PI * i as f32 / (frame_size - 1) as f32).cos());
frame[i] = signal[frame_start + i] * window;
}
}
let frame_slice =
frame
.as_slice()
.ok_or_else(|| EvaluationError::AudioProcessingError {
message: "Failed to get frame slice".to_string(),
source: None,
})?;
fft.process(frame_slice, &mut spectrum);
let power_spectrum: Vec<f32> = spectrum
.iter()
.map(scirs2_core::Complex::norm_sqr)
.collect();
for (bark_idx, &bark_freq) in self.bark_mapping.iter().enumerate() {
let bin_idx = (bark_freq * frame_size as f32 / self.sample_rate as f32) as usize;
if bin_idx < power_spectrum.len() {
let power_val = power_spectrum[bin_idx].max(0.0);
bark_spectrum[[frame_idx, bark_idx]] = power_val.sqrt();
}
}
}
Ok(bark_spectrum)
}
fn cognitive_modeling(
&self,
reference: &Array2<f32>,
degraded: &Array2<f32>,
) -> Result<Array1<f32>, EvaluationError> {
let (num_frames, num_bands) = reference.dim();
let mut disturbance_frames = Array1::zeros(num_frames);
for frame_idx in 0..num_frames {
let mut frame_disturbance = 0.0;
for band_idx in 0..num_bands {
let ref_val = reference[[frame_idx, band_idx]];
let deg_val = degraded[[frame_idx, band_idx]];
let weight = self.perceptual_weights[band_idx];
let ref_loudness = self.intensity_to_loudness(ref_val);
let deg_loudness = self.intensity_to_loudness(deg_val);
let difference = deg_loudness - ref_loudness;
let disturbance = if difference > 0.0 {
difference * weight
} else {
difference.abs() * weight * 0.5 };
frame_disturbance += disturbance;
}
disturbance_frames[frame_idx] = frame_disturbance;
}
Ok(disturbance_frames)
}
fn intensity_to_loudness(&self, intensity: f32) -> f32 {
if intensity <= 0.0 {
0.0
} else {
let db = 20.0 * intensity.log10();
if db < 0.0 {
0.0
} else {
(db / 40.0).powf(0.6)
}
}
}
fn calculate_final_score(
&self,
disturbance_frames: &Array1<f32>,
) -> Result<f32, EvaluationError> {
if disturbance_frames.is_empty() {
return Ok(1.0); }
let mut sorted_disturbances: Vec<f32> = disturbance_frames.to_vec();
sorted_disturbances.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let valid_disturbances: Vec<f32> = sorted_disturbances
.iter()
.filter(|&&x| x.is_finite() && !x.is_nan())
.copied()
.collect();
if valid_disturbances.is_empty() {
return Ok(4.5);
}
let d_95 = crate::precision::precise_percentile(valid_disturbances.clone(), 95.0) as f32;
let d_90 = crate::precision::precise_percentile(valid_disturbances.clone(), 90.0) as f32;
let d_80 = crate::precision::precise_percentile(valid_disturbances.clone(), 80.0) as f32;
let d_mean = valid_disturbances.iter().sum::<f32>() / valid_disturbances.len() as f32;
let d_indicator = if self.sample_rate == 8000 {
0.5 * d_95 + 0.3 * d_90 + 0.15 * d_80 + 0.05 * d_mean
} else {
0.4 * d_95 + 0.3 * d_90 + 0.2 * d_80 + 0.1 * d_mean
};
#[cfg(debug_assertions)]
{
eprintln!("PESQ Debug: disturbance frames: {}, d_indicator: {:.6}, d_95: {:.6}, d_mean: {:.6}",
valid_disturbances.len(), d_indicator, d_95, d_mean);
}
let pesq_score = if self.sample_rate == 8000 {
self.calculate_nb_pesq_score(d_indicator)
} else {
self.calculate_wb_pesq_score(d_indicator)
};
let calibrated_score = self.apply_perceptual_calibration(pesq_score);
#[cfg(debug_assertions)]
{
eprintln!(
"PESQ Debug: raw_score: {:.6}, calibrated: {:.6}",
pesq_score, calibrated_score
);
}
let clamped_score = calibrated_score.max(-0.5).min(4.5);
Ok(clamped_score)
}
fn calculate_nb_pesq_score(&self, d_indicator: f32) -> f32 {
if d_indicator < 0.1 {
4.5 - d_indicator * 5.0 } else if d_indicator < 0.5 {
4.0 - (d_indicator - 0.1) * 3.75 } else if d_indicator < 1.0 {
2.5 - (d_indicator - 0.5) * 2.5 } else if d_indicator < 2.0 {
1.25 - (d_indicator - 1.0) * 1.25 } else {
0.0 - (d_indicator - 2.0) * 0.25 }
}
fn calculate_wb_pesq_score(&self, d_indicator: f32) -> f32 {
if d_indicator < 0.1 {
4.5 - d_indicator * 4.0
} else if d_indicator < 0.4 {
4.1 - (d_indicator - 0.1) * 3.33
} else if d_indicator < 0.8 {
3.1 - (d_indicator - 0.4) * 2.5
} else if d_indicator < 1.5 {
2.1 - (d_indicator - 0.8) * 1.43
} else {
1.1 - (d_indicator - 1.5) * 0.73
}
}
fn apply_perceptual_calibration(&self, raw_score: f32) -> f32 {
if raw_score >= 4.0 {
return raw_score;
}
let x = (raw_score + 0.5) / 5.0;
let alpha = 2.5; let beta = 0.5;
let sigmoid = 1.0 / (1.0 + (-(alpha * (x - beta))).exp());
let calibrated = sigmoid * 5.0 - 0.5;
if calibrated > 3.5 {
calibrated + (calibrated - 3.5) * 0.2 } else if calibrated < 1.5 {
calibrated - (1.5 - calibrated) * 0.1 } else {
calibrated
}
}
fn create_bark_mapping(sample_rate: u32) -> Vec<f32> {
let num_bands = if sample_rate == 8000 { 18 } else { 24 };
let max_freq = sample_rate as f32 / 2.0;
(0..num_bands)
.map(|i| {
i as f32 * max_freq / (num_bands - 1) as f32
})
.collect()
}
fn create_perceptual_weights(bark_mapping: &[f32]) -> Array1<f32> {
let weights: Vec<f32> = bark_mapping
.iter()
.map(|&freq| {
if freq < 1000.0 {
0.5 + freq / 2000.0
} else if freq < 3000.0 {
1.0
} else {
1.0 - (freq - 3000.0) / 5000.0
}
.max(0.1)
.min(1.0)
})
.collect();
Array1::from_vec(weights)
}
}
#[cfg(test)]
mod tests {
use super::*;
use voirs_sdk::AudioBuffer;
#[tokio::test]
async fn test_pesq_evaluator_creation() {
let nb_evaluator = PESQEvaluator::new_narrowband().unwrap();
assert_eq!(nb_evaluator.sample_rate, 8000);
let wb_evaluator = PESQEvaluator::new_wideband().unwrap();
assert_eq!(wb_evaluator.sample_rate, 16000);
}
#[tokio::test]
async fn test_pesq_calculation() {
let evaluator = PESQEvaluator::new_narrowband().unwrap();
let duration_samples = 8 * 8000; let reference = AudioBuffer::new(vec![0.1; duration_samples], 8000, 1);
let degraded = AudioBuffer::new(vec![0.08; duration_samples], 8000, 1);
let pesq_score = evaluator
.calculate_pesq(&reference, °raded)
.await
.unwrap();
assert!(pesq_score >= -0.5);
assert!(pesq_score <= 4.5);
}
#[tokio::test]
async fn test_level_alignment() {
let evaluator = PESQEvaluator::new_narrowband().unwrap();
let reference = AudioBuffer::new(vec![0.1; 1000], 8000, 1);
let degraded = AudioBuffer::new(vec![0.2; 1000], 8000, 1);
let (ref_aligned, deg_aligned) = evaluator.level_alignment(&reference, °raded).unwrap();
let ref_mean_sq = ref_aligned.mapv(|x| x * x).mean().unwrap_or(0.0);
let deg_mean_sq = deg_aligned.mapv(|x| x * x).mean().unwrap_or(0.0);
let ref_rms = ref_mean_sq.max(0.0).sqrt();
let deg_rms = deg_mean_sq.max(0.0).sqrt();
assert!(
(ref_rms - deg_rms).abs() < 0.01,
"RMS levels should be similar after alignment: ref={}, deg={}",
ref_rms,
deg_rms
);
}
#[test]
fn test_correlation_calculation() {
let evaluator = PESQEvaluator::new_narrowband().unwrap();
let signal1 = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let signal2 = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
let correlation = evaluator.calculate_correlation(&signal1.view(), &signal2.view());
assert!((correlation - 1.0).abs() < 0.001);
let signal3 = Array1::from_vec(vec![4.0, 3.0, 2.0, 1.0]);
let correlation_neg = evaluator.calculate_correlation(&signal1.view(), &signal3.view());
assert!((correlation_neg + 1.0).abs() < 0.001); }
}