1use crate::fft::{Complex, Fft};
4use crate::Audio;
5use std::f64::consts::PI;
6
7#[derive(Clone, Debug, PartialEq)]
15pub struct ArtifactReport {
16 pub musical_noise_score: f64,
19 pub pumping_score: f64,
21 pub transient_loss_score: f64,
24 pub phase_distortion_score: Option<f64>,
26}
27
28#[derive(Clone, Debug)]
29pub struct BenchmarkReport {
30 pub frames: usize,
31 pub sample_rate: u32,
32 pub channels: usize,
33 pub si_sdr_db: f64,
34 pub si_snr_db: f64,
35 pub snr_db: f64,
36 pub segmental_snr_db: f64,
37 pub stereo_side_sdr_db: Option<f64>,
38 pub correlation_error: Option<f64>,
39 pub artifact_scores: ArtifactReport,
41 pub stoi: Option<f64>,
43 pub pesq: Option<f64>,
45 pub visqol: Option<f64>,
47 pub elapsed_ms: Option<f64>,
48 pub peak_rss_bytes: Option<u64>,
49}
50
51impl BenchmarkReport {
52 pub fn compare(reference: &Audio, test: &Audio) -> Result<Self, String> {
53 validate_pair(reference, test)?;
54 let frames = reference.frames();
55 let r = downmix(reference, frames);
56 let t = downmix(test, frames);
57 let artifact_scores = ArtifactReport::compare(reference, test)?;
58 let quality_metrics = crate::quality::QualityMetrics::compare(reference, test);
59 let (side_sdr, correlation_error) = if reference.channels.len() == 2 {
60 let rs = side(reference, frames);
61 let ts = side(test, frames);
62 (
63 Some(finite_db(si_sdr(&rs, &ts))),
64 Some(
65 (correlation(
66 &reference.channels[0][..frames],
67 &reference.channels[1][..frames],
68 ) - correlation(&test.channels[0][..frames], &test.channels[1][..frames]))
69 .abs(),
70 )
71 .filter(|value| value.is_finite()),
72 )
73 } else {
74 (None, None)
75 };
76 Ok(Self {
77 frames,
78 sample_rate: reference.sample_rate,
79 channels: reference.channels.len(),
80 si_sdr_db: finite_db(si_sdr(&r, &t)),
81 si_snr_db: finite_db(si_snr(&r, &t)),
82 snr_db: finite_db(snr(&r, &t)),
83 segmental_snr_db: finite_db(segmental_snr(&r, &t, reference.sample_rate)),
84 stereo_side_sdr_db: side_sdr,
85 correlation_error,
86 artifact_scores,
87 stoi: quality_metrics.stoi,
88 pesq: quality_metrics.pesq,
89 visqol: quality_metrics.visqol,
90 elapsed_ms: None,
91 peak_rss_bytes: None,
92 })
93 }
94
95 pub fn json(&self) -> String {
96 format!("{{\"frames\":{},\"sample_rate\":{},\"channels\":{},\"si_sdr_db\":{},\"si_snr_db\":{},\"snr_db\":{},\"segmental_snr_db\":{},\"stereo_side_sdr_db\":{},\"correlation_error\":{},\"artifact_scores\":{},\"stoi\":{},\"pesq\":{},\"visqol\":{},\"elapsed_ms\":{},\"peak_rss_bytes\":{}}}", self.frames, self.sample_rate, self.channels, json_number(self.si_sdr_db), json_number(self.si_snr_db), json_number(self.snr_db), json_number(self.segmental_snr_db), optional(self.stereo_side_sdr_db), optional(self.correlation_error), self.artifact_scores.json(), optional(self.stoi), optional(self.pesq), optional(self.visqol), optional(self.elapsed_ms), self.peak_rss_bytes.map_or_else(|| "null".into(), |v| v.to_string()))
97 }
98
99 pub fn markdown(&self) -> String {
100 format!("| Metric | Value |\n|---|---:|\n| SI-SDR | {:.3} dB |\n| SI-SNR | {:.3} dB |\n| SNR | {:.3} dB |\n| Segmental SNR | {:.3} dB |\n| Stereo side SDR | {} |\n| Correlation error | {} |\n| Musical-noise score (0=none) | {:.3} |\n| Pumping score (0=none) | {:.3} |\n| Transient-loss score (0=none) | {:.3} |\n| Phase-distortion score (0=none) | {} |\n| STOI (-1–1, higher is better) | {} |\n| PESQ (licensed adapter required) | {} |\n| ViSQOL MOS-LQO (1–5) | {} |", self.si_sdr_db, self.si_snr_db, self.snr_db, self.segmental_snr_db, db(self.stereo_side_sdr_db), display(self.correlation_error, 6), self.artifact_scores.musical_noise_score, self.artifact_scores.pumping_score, self.artifact_scores.transient_loss_score, display(self.artifact_scores.phase_distortion_score, 3), display(self.stoi, 4), display(self.pesq, 3), display(self.visqol, 3))
101 }
102}
103
104impl ArtifactReport {
105 pub fn compare(reference: &Audio, test: &Audio) -> Result<Self, String> {
108 validate_pair(reference, test)?;
109 let frames = reference.frames();
110 let ref_mix = downmix(reference, frames);
111 let test_mix = downmix(test, frames);
112 let observations = collect_artifact_observations(reference, test, &ref_mix, &test_mix);
113
114 let musical_noise_score = normalized_score(
115 observations.musical_noise_excess,
116 observations.test_spectral_energy,
117 );
118 let pumping_score = pumping_score(&observations.rms_pairs);
119 let transient_loss_score =
120 transient_loss_score(&observations.ref_flux, &observations.test_flux);
121 let phase_distortion_score = if reference.channels.len() == 2 {
122 Some(normalized_score(
123 observations.phase_error,
124 observations.phase_weight,
125 ))
126 } else {
127 None
128 };
129
130 Ok(Self {
131 musical_noise_score,
132 pumping_score,
133 transient_loss_score,
134 phase_distortion_score,
135 })
136 }
137
138 pub fn json(&self) -> String {
140 format!(
141 "{{\"musical_noise_score\":{},\"pumping_score\":{},\"transient_loss_score\":{},\"phase_distortion_score\":{}}}",
142 json_number(self.musical_noise_score),
143 json_number(self.pumping_score),
144 json_number(self.transient_loss_score),
145 optional(self.phase_distortion_score),
146 )
147 }
148}
149
150#[derive(Clone, Debug)]
151pub struct ComparisonReport {
152 pub noisy: BenchmarkReport,
153 pub enhanced: BenchmarkReport,
154}
155
156impl ComparisonReport {
157 pub fn compare(clean: &Audio, noisy: &Audio, enhanced: &Audio) -> Result<Self, String> {
158 validate_pair(clean, noisy).map_err(|error| format!("noisy comparison: {error}"))?;
159 validate_pair(clean, enhanced).map_err(|error| format!("enhanced comparison: {error}"))?;
160 Ok(Self {
161 noisy: BenchmarkReport::compare(clean, noisy)?,
162 enhanced: BenchmarkReport::compare(clean, enhanced)?,
163 })
164 }
165
166 pub fn json(&self) -> String {
167 format!(
168 "{{\"noisy\":{},\"enhanced\":{},\"improvement\":{{\"si_sdr_db\":{},\"si_snr_db\":{},\"snr_db\":{},\"segmental_snr_db\":{},\"stereo_side_sdr_db\":{},\"correlation_error\":{},\"stoi\":{},\"pesq\":{},\"visqol\":{},\"musical_noise_score\":{},\"pumping_score\":{},\"transient_loss_score\":{},\"phase_distortion_score\":{}}}}}",
169 self.noisy.json(), self.enhanced.json(),
170 json_number(self.enhanced.si_sdr_db - self.noisy.si_sdr_db),
171 json_number(self.enhanced.si_snr_db - self.noisy.si_snr_db),
172 json_number(self.enhanced.snr_db - self.noisy.snr_db),
173 json_number(self.enhanced.segmental_snr_db - self.noisy.segmental_snr_db),
174 optional_difference(self.enhanced.stereo_side_sdr_db, self.noisy.stereo_side_sdr_db),
175 optional_difference(self.noisy.correlation_error, self.enhanced.correlation_error),
176 optional_difference(self.enhanced.stoi, self.noisy.stoi),
177 optional_difference(self.enhanced.pesq, self.noisy.pesq),
178 optional_difference(self.enhanced.visqol, self.noisy.visqol),
179 json_number(self.noisy.artifact_scores.musical_noise_score - self.enhanced.artifact_scores.musical_noise_score),
180 json_number(self.noisy.artifact_scores.pumping_score - self.enhanced.artifact_scores.pumping_score),
181 json_number(self.noisy.artifact_scores.transient_loss_score - self.enhanced.artifact_scores.transient_loss_score),
182 optional_difference(self.noisy.artifact_scores.phase_distortion_score, self.enhanced.artifact_scores.phase_distortion_score),
183 )
184 }
185
186 pub fn markdown(&self) -> String {
187 format!(
188 "| Metric | Noisy | Enhanced | Improvement |\n|---|---:|---:|---:|\n| SI-SDR | {:.3} dB | {:.3} dB | {:+.3} dB |\n| SI-SNR | {:.3} dB | {:.3} dB | {:+.3} dB |\n| SNR | {:.3} dB | {:.3} dB | {:+.3} dB |\n| Segmental SNR | {:.3} dB | {:.3} dB | {:+.3} dB |\n| Stereo side SDR (higher is better) | {} | {} | {} |\n| Correlation error (lower is better) | {} | {} | {} |\n| STOI (higher is better) | {} | {} | {} |\n| PESQ (higher is better; licensed adapter) | {} | {} | {} |\n| ViSQOL MOS-LQO (higher is better) | {} | {} | {} |\n| Musical-noise score (lower is better) | {:.3} | {:.3} | {:+.3} |\n| Pumping score (lower is better) | {:.3} | {:.3} | {:+.3} |\n| Transient-loss score (lower is better) | {:.3} | {:.3} | {:+.3} |\n| Phase-distortion score (lower is better) | {} | {} | {} |\n\nArtifact scores are deterministic screening indicators in [0, 1], not perceptual listening-test scores. STOI is implemented natively. ViSQOL is measured when the `visqol` feature is enabled. PESQ remains unavailable because its ITU-T reference implementation requires a separately licensed external adapter.",
189 self.noisy.si_sdr_db, self.enhanced.si_sdr_db, self.enhanced.si_sdr_db - self.noisy.si_sdr_db,
190 self.noisy.si_snr_db, self.enhanced.si_snr_db, self.enhanced.si_snr_db - self.noisy.si_snr_db,
191 self.noisy.snr_db, self.enhanced.snr_db, self.enhanced.snr_db - self.noisy.snr_db,
192 self.noisy.segmental_snr_db, self.enhanced.segmental_snr_db, self.enhanced.segmental_snr_db - self.noisy.segmental_snr_db,
193 db(self.noisy.stereo_side_sdr_db), db(self.enhanced.stereo_side_sdr_db), db(optional_difference_value(self.enhanced.stereo_side_sdr_db, self.noisy.stereo_side_sdr_db)),
194 display(self.noisy.correlation_error, 6), display(self.enhanced.correlation_error, 6), display(optional_difference_value(self.noisy.correlation_error, self.enhanced.correlation_error), 6),
195 display(self.noisy.stoi, 4), display(self.enhanced.stoi, 4), display(optional_difference_value(self.enhanced.stoi, self.noisy.stoi), 4),
196 display(self.noisy.pesq, 3), display(self.enhanced.pesq, 3), display(optional_difference_value(self.enhanced.pesq, self.noisy.pesq), 3),
197 display(self.noisy.visqol, 3), display(self.enhanced.visqol, 3), display(optional_difference_value(self.enhanced.visqol, self.noisy.visqol), 3),
198 self.noisy.artifact_scores.musical_noise_score, self.enhanced.artifact_scores.musical_noise_score, self.noisy.artifact_scores.musical_noise_score - self.enhanced.artifact_scores.musical_noise_score,
199 self.noisy.artifact_scores.pumping_score, self.enhanced.artifact_scores.pumping_score, self.noisy.artifact_scores.pumping_score - self.enhanced.artifact_scores.pumping_score,
200 self.noisy.artifact_scores.transient_loss_score, self.enhanced.artifact_scores.transient_loss_score, self.noisy.artifact_scores.transient_loss_score - self.enhanced.artifact_scores.transient_loss_score,
201 display(self.noisy.artifact_scores.phase_distortion_score, 3), display(self.enhanced.artifact_scores.phase_distortion_score, 3), display(optional_difference_value(self.noisy.artifact_scores.phase_distortion_score, self.enhanced.artifact_scores.phase_distortion_score), 3),
202 )
203 }
204
205 pub fn html(&self) -> String {
206 let rows = self
207 .markdown()
208 .lines()
209 .skip(2)
210 .take_while(|line| !line.is_empty())
215 .map(|line| {
216 let cells = line
217 .trim_matches('|')
218 .split('|')
219 .map(str::trim)
220 .map(|cell| format!("<td>{cell}</td>"))
221 .collect::<String>();
222 format!("<tr>{cells}</tr>")
223 })
224 .collect::<String>();
225 format!("<!doctype html><meta charset=\"utf-8\"><title>denoize comparison</title><style>body{{font-family:system-ui;max-width:900px;margin:3rem auto}}table{{border-collapse:collapse}}td,th{{padding:.5rem 1rem;border:1px solid #ccc;text-align:right}}td:first-child{{text-align:left}}</style><h1>denoize quality comparison</h1><table><thead><tr><th>Metric</th><th>Noisy</th><th>Enhanced</th><th>Improvement</th></tr></thead><tbody>{rows}</tbody></table><p>Artifact scores are deterministic screening indicators in [0, 1], lower is better. STOI is native; ViSQOL requires the <code>visqol</code> feature; PESQ requires a separately licensed external adapter.</p>")
226 }
227}
228
229fn json_number(value: f64) -> String {
230 if value.is_finite() {
231 format!("{value:.6}")
232 } else {
233 "null".into()
234 }
235}
236
237fn finite_db(value: f64) -> f64 {
238 if value.is_finite() {
239 value
240 } else {
241 -120.0
242 }
243}
244
245fn optional(v: Option<f64>) -> String {
246 v.map_or_else(|| "null".into(), json_number)
247}
248fn display(v: Option<f64>, precision: usize) -> String {
249 v.filter(|value| value.is_finite())
250 .map_or_else(|| "n/a".into(), |v| format!("{v:.precision$}"))
251}
252fn db(v: Option<f64>) -> String {
253 v.filter(|value| value.is_finite())
254 .map_or_else(|| "n/a".into(), |v| format!("{v:.3} dB"))
255}
256
257fn optional_difference(a: Option<f64>, b: Option<f64>) -> String {
258 optional(optional_difference_value(a, b))
259}
260
261fn optional_difference_value(a: Option<f64>, b: Option<f64>) -> Option<f64> {
262 match (a, b) {
263 (Some(a), Some(b)) if a.is_finite() && b.is_finite() => {
264 let difference = a - b;
265 difference.is_finite().then_some(difference)
266 }
267 _ => None,
268 }
269}
270
271fn validate_pair(reference: &Audio, test: &Audio) -> Result<(), String> {
272 let reference_frames = validate_audio_shape(reference, "reference")?;
273 let test_frames = validate_audio_shape(test, "test")?;
274 if reference.sample_rate != test.sample_rate {
275 return Err(format!(
276 "benchmark sample rates differ: reference is {} Hz, test is {} Hz",
277 reference.sample_rate, test.sample_rate
278 ));
279 }
280 if reference.channels.len() != test.channels.len() {
281 return Err(format!(
282 "benchmark channel counts differ: reference has {}, test has {}",
283 reference.channels.len(),
284 test.channels.len()
285 ));
286 }
287 if reference_frames != test_frames {
288 return Err(format!(
289 "benchmark frame counts differ: reference has {reference_frames}, test has {test_frames}"
290 ));
291 }
292 Ok(())
293}
294
295fn validate_audio_shape(audio: &Audio, name: &str) -> Result<usize, String> {
296 if audio.sample_rate == 0 {
297 return Err(format!("benchmark {name} sample rate is zero"));
298 }
299 let Some(first) = audio.channels.first() else {
300 return Err(format!("benchmark {name} has no channels"));
301 };
302 let frames = first.len();
303 if let Some((channel, actual)) = audio
304 .channels
305 .iter()
306 .enumerate()
307 .skip(1)
308 .map(|(channel, samples)| (channel, samples.len()))
309 .find(|(_, actual)| *actual != frames)
310 {
311 return Err(format!(
312 "benchmark {name} channels have inconsistent frame counts: channel 0 has {frames}, channel {channel} has {actual}"
313 ));
314 }
315 if frames == 0 {
316 return Err(format!("benchmark {name} has no frames"));
317 }
318 Ok(frames)
319}
320
321#[derive(Default)]
322struct ArtifactObservations {
323 musical_noise_excess: f64,
324 test_spectral_energy: f64,
325 rms_pairs: Vec<(f64, f64)>,
326 ref_flux: Vec<f64>,
327 test_flux: Vec<f64>,
328 phase_error: f64,
329 phase_weight: f64,
330}
331
332fn collect_artifact_observations(
336 reference: &Audio,
337 test: &Audio,
338 ref_mix: &[f64],
339 test_mix: &[f64],
340) -> ArtifactObservations {
341 let frames = ref_mix.len().min(test_mix.len());
342 let frame_size = artifact_frame_size(frames);
343 let hop = frame_size / 2;
344 let window = hann_window(frame_size);
345 let fft = Fft::new(frame_size);
346 let nbins = fft.nbins();
347 let starts = frame_starts(frames, frame_size, hop);
348
349 let mut ref_buffer = vec![Complex::default(); frame_size];
350 let mut test_buffer = vec![Complex::default(); frame_size];
351 let mut ref_left_buffer = vec![Complex::default(); frame_size];
352 let mut ref_right_buffer = vec![Complex::default(); frame_size];
353 let mut test_left_buffer = vec![Complex::default(); frame_size];
354 let mut test_right_buffer = vec![Complex::default(); frame_size];
355 let mut ref_magnitude = vec![0.0; nbins];
356 let mut test_magnitude = vec![0.0; nbins];
357 let mut prev_ref_magnitude = vec![0.0; nbins];
358 let mut prev_test_magnitude = vec![0.0; nbins];
359 let stereo = reference.channels.len() == 2;
360 let mut observations = ArtifactObservations {
361 rms_pairs: Vec::with_capacity(starts.len()),
362 ref_flux: Vec::with_capacity(starts.len()),
363 test_flux: Vec::with_capacity(starts.len()),
364 ..ArtifactObservations::default()
365 };
366
367 for (frame_index, &start) in starts.iter().enumerate() {
368 fill_windowed(&mut ref_buffer, ref_mix, start, &window);
369 fill_windowed(&mut test_buffer, test_mix, start, &window);
370 fft.forward(&mut ref_buffer);
371 fft.forward(&mut test_buffer);
372 magnitudes(&ref_buffer, &mut ref_magnitude);
373 magnitudes(&test_buffer, &mut test_magnitude);
374
375 observations.ref_flux.push(if frame_index == 0 {
376 0.0
377 } else {
378 spectral_flux(&ref_magnitude, &prev_ref_magnitude)
379 });
380 observations.test_flux.push(if frame_index == 0 {
381 0.0
382 } else {
383 spectral_flux(&test_magnitude, &prev_test_magnitude)
384 });
385 prev_ref_magnitude.copy_from_slice(&ref_magnitude);
386 prev_test_magnitude.copy_from_slice(&test_magnitude);
387
388 let ref_rms = frame_rms(ref_mix, start, frame_size);
389 let test_rms = frame_rms(test_mix, start, frame_size);
390 observations.rms_pairs.push((ref_rms, test_rms));
391
392 for k in 1..nbins.saturating_sub(1) {
395 let test_power = test_magnitude[k] * test_magnitude[k];
396 let ref_power = ref_magnitude[k] * ref_magnitude[k];
397 let excess = (test_power - 1.15 * ref_power).max(0.0);
398 let neighbour_power = 0.5
399 * (test_magnitude[k - 1] * test_magnitude[k - 1]
400 + test_magnitude[k + 1] * test_magnitude[k + 1]);
401 let prominence = ((test_power / (neighbour_power + 1e-30)) - 1.0).clamp(0.0, 4.0) / 4.0;
402 observations.musical_noise_excess += excess * prominence;
403 }
404 observations.test_spectral_energy += test_magnitude.iter().map(|m| m * m).sum::<f64>();
405
406 if stereo {
407 fill_windowed(&mut ref_left_buffer, &reference.channels[0], start, &window);
408 fill_windowed(
409 &mut ref_right_buffer,
410 &reference.channels[1],
411 start,
412 &window,
413 );
414 fill_windowed(&mut test_left_buffer, &test.channels[0], start, &window);
415 fill_windowed(&mut test_right_buffer, &test.channels[1], start, &window);
416 fft.forward(&mut ref_left_buffer);
417 fft.forward(&mut ref_right_buffer);
418 fft.forward(&mut test_left_buffer);
419 fft.forward(&mut test_right_buffer);
420 let (error, weight) = phase_error_for_frame(
421 &ref_left_buffer,
422 &ref_right_buffer,
423 &test_left_buffer,
424 &test_right_buffer,
425 );
426 observations.phase_error += error;
427 observations.phase_weight += weight;
428 }
429 }
430
431 observations
432}
433
434fn artifact_frame_size(frames: usize) -> usize {
435 frames.min(1024).max(64).next_power_of_two().min(1024)
438}
439
440fn frame_starts(frames: usize, frame_size: usize, hop: usize) -> Vec<usize> {
441 let mut starts = Vec::with_capacity((frames / hop).saturating_add(1));
442 let mut start = 0;
443 while start < frames {
444 starts.push(start);
445 if start >= frames.saturating_sub(frame_size) {
446 break;
447 }
448 start += hop;
449 }
450 starts
451}
452
453fn hann_window(size: usize) -> Vec<f64> {
454 (0..size)
455 .map(|index| 0.5 - 0.5 * (2.0 * PI * index as f64 / (size.saturating_sub(1) as f64)).cos())
456 .collect()
457}
458
459fn fill_windowed(buffer: &mut [Complex], signal: &[f64], start: usize, window: &[f64]) {
460 for (index, slot) in buffer.iter_mut().enumerate() {
461 let sample = signal
462 .get(start + index)
463 .copied()
464 .filter(|sample| sample.is_finite())
465 .unwrap_or(0.0);
466 *slot = Complex::new(sample * window[index], 0.0);
467 }
468}
469
470fn magnitudes(spectrum: &[Complex], output: &mut [f64]) {
471 for (index, magnitude) in output.iter_mut().enumerate() {
472 let value = spectrum[index];
473 *magnitude = value.re.hypot(value.im);
474 }
475}
476
477fn spectral_flux(current: &[f64], previous: &[f64]) -> f64 {
478 let rise = current
479 .iter()
480 .zip(previous)
481 .map(|(current, previous)| (current - previous).max(0.0))
482 .sum::<f64>();
483 let previous_energy = previous.iter().sum::<f64>();
484 (rise / (previous_energy + 1e-12)).clamp(0.0, 100.0)
485}
486
487fn frame_rms(signal: &[f64], start: usize, frame_size: usize) -> f64 {
488 let end = (start + frame_size).min(signal.len());
489 if start >= end {
490 return 0.0;
491 }
492 let sum = signal[start..end]
493 .iter()
494 .filter(|sample| sample.is_finite())
495 .map(|sample| sample * sample)
496 .sum::<f64>();
497 (sum / (end - start) as f64).sqrt()
498}
499
500fn phase_error_for_frame(
501 reference_left: &[Complex],
502 reference_right: &[Complex],
503 test_left: &[Complex],
504 test_right: &[Complex],
505) -> (f64, f64) {
506 let nbins = reference_left.len() / 2 + 1;
507 let mut error = 0.0;
508 let mut weight = 0.0;
509 for k in 1..nbins.saturating_sub(1) {
510 let reference_cross = cross_spectrum(reference_left[k], reference_right[k]);
511 let test_cross = cross_spectrum(test_left[k], test_right[k]);
512 let reference_magnitude = reference_cross.re.hypot(reference_cross.im);
513 let test_magnitude = test_cross.re.hypot(test_cross.im);
514 if reference_magnitude <= 1e-12 || test_magnitude <= 1e-12 {
515 continue;
516 }
517 let cosine = ((reference_cross.re * test_cross.re + reference_cross.im * test_cross.im)
518 / (reference_magnitude * test_magnitude))
519 .clamp(-1.0, 1.0);
520 let bin_weight = reference_magnitude.min(test_magnitude);
521 error += 0.5 * (1.0 - cosine) * bin_weight;
522 weight += bin_weight;
523 }
524 (error, weight)
525}
526
527fn cross_spectrum(left: Complex, right: Complex) -> Complex {
528 Complex::new(
530 left.re * right.re + left.im * right.im,
531 left.im * right.re - left.re * right.im,
532 )
533}
534
535fn normalized_score(numerator: f64, denominator: f64) -> f64 {
536 if denominator <= 1e-30 || !numerator.is_finite() || !denominator.is_finite() {
537 0.0
538 } else {
539 (numerator / denominator).clamp(0.0, 1.0)
540 }
541}
542
543fn pumping_score(rms_pairs: &[(f64, f64)]) -> f64 {
544 let reference_peak = rms_pairs
545 .iter()
546 .map(|(reference, _)| *reference)
547 .fold(0.0, f64::max);
548 if reference_peak <= 1e-9 {
549 return 0.0;
550 }
551 let floor = reference_peak * 0.01;
552 let mut previous_gain: Option<f64> = None;
553 let mut total_change = 0.0;
554 let mut count = 0usize;
555 for &(reference, test) in rms_pairs {
556 if reference <= floor {
557 previous_gain = None;
558 continue;
559 }
560 let gain_db = (20.0 * ((test + floor) / (reference + floor)).log10()).clamp(-60.0, 60.0);
561 if let Some(previous) = previous_gain {
562 total_change += (gain_db - previous).abs();
563 count += 1;
564 }
565 previous_gain = Some(gain_db);
566 }
567 if count == 0 {
568 0.0
569 } else {
570 (total_change / count as f64 / 12.0).clamp(0.0, 1.0)
571 }
572}
573
574fn transient_loss_score(reference_flux: &[f64], test_flux: &[f64]) -> f64 {
575 let maximum = reference_flux.iter().copied().fold(0.0, f64::max);
576 let threshold = (maximum * 0.15).max(0.05);
577 let mut weighted_loss = 0.0;
578 let mut total_weight = 0.0;
579 for (&reference, &test) in reference_flux.iter().zip(test_flux) {
580 if reference < threshold {
581 continue;
582 }
583 weighted_loss += ((reference - test).max(0.0) / reference.max(1e-12)) * reference;
584 total_weight += reference;
585 }
586 normalized_score(weighted_loss, total_weight)
587}
588
589fn downmix(a: &Audio, n: usize) -> Vec<f64> {
590 (0..n)
591 .map(|i| a.channels.iter().map(|c| c[i]).sum::<f64>() / a.channels.len() as f64)
592 .collect()
593}
594fn side(a: &Audio, n: usize) -> Vec<f64> {
595 (0..n)
596 .map(|i| (a.channels[0][i] - a.channels[1][i]) * 0.5)
597 .collect()
598}
599
600pub fn si_sdr(reference: &[f64], estimate: &[f64]) -> f64 {
601 let dot = reference
602 .iter()
603 .zip(estimate)
604 .map(|(a, b)| a * b)
605 .sum::<f64>();
606 let scale = dot / reference.iter().map(|x| x * x).sum::<f64>().max(1e-30);
607 let target_energy = reference.iter().map(|x| (x * scale).powi(2)).sum::<f64>();
608 let noise_energy = reference
609 .iter()
610 .zip(estimate)
611 .map(|(a, b)| (a * scale - b).powi(2))
612 .sum::<f64>();
613 10.0 * (target_energy / noise_energy.max(1e-30)).log10()
614}
615
616pub fn si_snr(reference: &[f64], estimate: &[f64]) -> f64 {
617 let rm = reference.iter().sum::<f64>() / reference.len() as f64;
618 let em = estimate.iter().sum::<f64>() / estimate.len() as f64;
619 si_sdr(
620 &reference.iter().map(|x| x - rm).collect::<Vec<_>>(),
621 &estimate.iter().map(|x| x - em).collect::<Vec<_>>(),
622 )
623}
624
625pub fn snr(reference: &[f64], estimate: &[f64]) -> f64 {
626 let signal = reference.iter().map(|sample| sample * sample).sum::<f64>();
627 let noise = reference
628 .iter()
629 .zip(estimate)
630 .map(|(a, b)| (a - b).powi(2))
631 .sum::<f64>();
632 10.0 * (signal / noise.max(1e-30)).log10()
633}
634
635pub fn segmental_snr(reference: &[f64], estimate: &[f64], sample_rate: u32) -> f64 {
636 let window = (sample_rate as usize / 50).max(1);
637 let mut values = Vec::new();
638 for (r, e) in reference.chunks(window).zip(estimate.chunks(window)) {
639 let signal = r.iter().map(|sample| sample * sample).sum::<f64>();
640 if signal > 1e-12 {
641 values.push(snr(r, e).clamp(-10.0, 35.0));
642 }
643 }
644 values.iter().sum::<f64>() / values.len().max(1) as f64
645}
646
647fn correlation(a: &[f64], b: &[f64]) -> f64 {
648 let dot = a.iter().zip(b).map(|(a, b)| a * b).sum::<f64>();
649 dot / (a.iter().map(|x| x * x).sum::<f64>() * b.iter().map(|x| x * x).sum::<f64>())
650 .sqrt()
651 .max(1e-30)
652}
653
654#[cfg(test)]
655mod tests {
656 use super::*;
657
658 #[test]
659 fn metrics_ignore_gain() {
660 let reference = [1.0, -1.0, 0.5, -0.5];
661 let estimate = [0.5, -0.5, 0.25, -0.25];
662 assert!(si_sdr(&reference, &estimate) > 250.0);
663 assert!(si_snr(&reference, &estimate) > 250.0);
664 }
665
666 #[test]
667 fn comparison_reports_quality_improvement_in_all_formats() {
668 let clean = Audio {
669 sample_rate: 16_000,
670 channels: vec![(0..1600)
671 .map(|index| (index as f64 * 0.031).sin() * 0.5)
672 .collect()],
673 bits_per_sample: 16,
674 sample_format: hound::SampleFormat::Int,
675 channel_mask: None,
676 };
677 let noisy = Audio {
678 sample_rate: clean.sample_rate,
679 channels: vec![clean.channels[0]
680 .iter()
681 .enumerate()
682 .map(|(index, sample)| sample + if index % 2 == 0 { 0.1 } else { -0.1 })
683 .collect()],
684 bits_per_sample: 16,
685 sample_format: hound::SampleFormat::Int,
686 channel_mask: None,
687 };
688 let enhanced = Audio {
689 sample_rate: clean.sample_rate,
690 channels: vec![clean.channels[0]
691 .iter()
692 .enumerate()
693 .map(|(index, sample)| sample + if index % 2 == 0 { 0.02 } else { -0.02 })
694 .collect()],
695 bits_per_sample: 16,
696 sample_format: hound::SampleFormat::Int,
697 channel_mask: None,
698 };
699
700 let report = ComparisonReport::compare(&clean, &noisy, &enhanced).unwrap();
701 assert!(report.enhanced.snr_db > report.noisy.snr_db);
702 assert!(report.json().contains("\"improvement\""));
703 assert!(report.json().contains("\"artifact_scores\""));
704 assert!(report.json().contains("\"stoi\""));
705 assert!(report.json().contains("\"pesq\""));
706 assert!(report.json().contains("\"stereo_side_sdr_db\""));
707 assert!(report.json().contains("\"correlation_error\""));
708 assert!(report.json().contains("\"visqol\""));
709 assert!(report.markdown().contains("Segmental SNR"));
710 assert!(report.markdown().contains("Stereo side SDR"));
711 assert!(report.markdown().contains("Correlation error"));
712 assert!(report.markdown().contains("STOI"));
713 assert!(report.markdown().contains("PESQ"));
714 assert!(report.markdown().contains("ViSQOL"));
715 assert!(report.markdown().contains("Musical-noise score"));
716 let html = report.html();
717 assert!(html.starts_with("<!doctype html>"));
718 assert!(html.contains("Transient-loss score"));
719 assert!(html.contains("Phase-distortion score"));
720 }
721
722 #[test]
723 fn silent_comparison_metrics_stay_finite_and_json_safe() {
724 let audio = mono(vec![0.0; 1600]);
725 let report = ComparisonReport::compare(&audio, &audio, &audio).unwrap();
726 for metrics in [&report.noisy, &report.enhanced] {
727 assert!(metrics.si_sdr_db.is_finite());
728 assert!(metrics.si_snr_db.is_finite());
729 assert!(metrics.snr_db.is_finite());
730 assert!(metrics.segmental_snr_db.is_finite());
731 }
732 let json = report.json();
733 assert!(!json.contains("NaN"));
734 assert!(!json.contains("inf"));
735 assert!(json.contains("\"si_sdr_db\":-120.000000"));
736 }
737
738 #[test]
739 fn rejects_truncated_benchmark_input() {
740 let reference = mono(vec![0.0; 1600]);
741 let test = mono(vec![0.0; 800]);
742 let error = BenchmarkReport::compare(&reference, &test).unwrap_err();
743 assert_eq!(
744 error,
745 "benchmark frame counts differ: reference has 1600, test has 800"
746 );
747 let artifact_error = ArtifactReport::compare(&reference, &test).unwrap_err();
748 assert_eq!(artifact_error, error);
749 }
750
751 #[test]
752 fn comparison_identifies_truncated_noisy_and_enhanced_inputs() {
753 let clean = mono(vec![0.0; 1600]);
754 let truncated = mono(vec![0.0; 800]);
755
756 let noisy_error = ComparisonReport::compare(&clean, &truncated, &clean).unwrap_err();
757 assert_eq!(
758 noisy_error,
759 "noisy comparison: benchmark frame counts differ: reference has 1600, test has 800"
760 );
761
762 let enhanced_error = ComparisonReport::compare(&clean, &clean, &truncated).unwrap_err();
763 assert_eq!(
764 enhanced_error,
765 "enhanced comparison: benchmark frame counts differ: reference has 1600, test has 800"
766 );
767 }
768
769 #[test]
770 fn rejects_ragged_channels_without_panicking() {
771 let reference = Audio {
772 sample_rate: 16_000,
773 channels: vec![vec![0.0; 1600], vec![0.0; 800]],
774 bits_per_sample: 32,
775 sample_format: hound::SampleFormat::Float,
776 channel_mask: None,
777 };
778 let test = stereo(vec![0.0; 1600], vec![0.0; 1600]);
779 let error = BenchmarkReport::compare(&reference, &test).unwrap_err();
780 assert_eq!(
781 error,
782 "benchmark reference channels have inconsistent frame counts: channel 0 has 1600, channel 1 has 800"
783 );
784
785 let valid = stereo(vec![0.0; 1600], vec![0.0; 1600]);
786 let ragged_test = Audio {
787 sample_rate: 16_000,
788 channels: vec![vec![0.0; 1600], vec![0.0; 800]],
789 bits_per_sample: 32,
790 sample_format: hound::SampleFormat::Float,
791 channel_mask: None,
792 };
793 let error = BenchmarkReport::compare(&valid, &ragged_test).unwrap_err();
794 assert_eq!(
795 error,
796 "benchmark test channels have inconsistent frame counts: channel 0 has 1600, channel 1 has 800"
797 );
798 }
799
800 #[test]
801 fn rejects_zero_sample_rate() {
802 let mut reference = mono(vec![0.0; 1600]);
803 reference.sample_rate = 0;
804 let test = mono(vec![0.0; 1600]);
805 let error = BenchmarkReport::compare(&reference, &test).unwrap_err();
806 assert_eq!(error, "benchmark reference sample rate is zero");
807 }
808
809 #[test]
810 fn rejects_empty_or_mismatched_benchmark_shapes() {
811 let valid = mono(vec![0.0; 1600]);
812 let no_channels = Audio {
813 sample_rate: 16_000,
814 channels: Vec::new(),
815 bits_per_sample: 32,
816 sample_format: hound::SampleFormat::Float,
817 channel_mask: None,
818 };
819 assert_eq!(
820 BenchmarkReport::compare(&no_channels, &valid).unwrap_err(),
821 "benchmark reference has no channels"
822 );
823
824 let no_frames = mono(Vec::new());
825 assert_eq!(
826 BenchmarkReport::compare(&valid, &no_frames).unwrap_err(),
827 "benchmark test has no frames"
828 );
829
830 let mut different_rate = valid.clone();
831 different_rate.sample_rate = 48_000;
832 assert_eq!(
833 BenchmarkReport::compare(&valid, &different_rate).unwrap_err(),
834 "benchmark sample rates differ: reference is 16000 Hz, test is 48000 Hz"
835 );
836
837 let stereo_test = stereo(vec![0.0; 1600], vec![0.0; 1600]);
838 assert_eq!(
839 BenchmarkReport::compare(&valid, &stereo_test).unwrap_err(),
840 "benchmark channel counts differ: reference has 1, test has 2"
841 );
842 }
843
844 fn mono(samples: Vec<f64>) -> Audio {
845 Audio {
846 sample_rate: 16_000,
847 channels: vec![samples],
848 bits_per_sample: 32,
849 sample_format: hound::SampleFormat::Float,
850 channel_mask: None,
851 }
852 }
853
854 fn stereo(left: Vec<f64>, right: Vec<f64>) -> Audio {
855 Audio {
856 sample_rate: 16_000,
857 channels: vec![left, right],
858 bits_per_sample: 32,
859 sample_format: hound::SampleFormat::Float,
860 channel_mask: None,
861 }
862 }
863
864 #[test]
865 fn identical_audio_has_no_artifact_scores() {
866 let signal = (0..16_384)
867 .map(|index| {
868 (2.0 * PI * 440.0 * index as f64 / 16_000.0).sin() * 0.4
869 + (2.0 * PI * 913.0 * index as f64 / 16_000.0).sin() * 0.1
870 })
871 .collect::<Vec<_>>();
872 let audio = mono(signal);
873 let report = ArtifactReport::compare(&audio, &audio).unwrap();
874 assert!(report.musical_noise_score < 1e-12);
875 assert!(report.pumping_score < 1e-12);
876 assert!(report.transient_loss_score < 1e-12);
877 assert_eq!(report.phase_distortion_score, None);
878 }
879
880 #[test]
881 fn detects_frame_level_pumping() {
882 let reference = (0..16_384)
883 .map(|index| (2.0 * PI * 440.0 * index as f64 / 16_000.0).sin() * 0.5)
884 .collect::<Vec<_>>();
885 let test = reference
886 .iter()
887 .enumerate()
888 .map(|(index, sample)| {
889 let gain = if (index / 2_048) % 2 == 0 { 1.0 } else { 0.25 };
890 sample * gain
891 })
892 .collect::<Vec<_>>();
893 let report = ArtifactReport::compare(&mono(reference), &mono(test)).unwrap();
894 assert!(
895 report.pumping_score > 0.2,
896 "pumping score: {}",
897 report.pumping_score
898 );
899 }
900
901 #[test]
902 fn detects_transient_loss() {
903 let mut reference = vec![0.0; 16_384];
904 for index in (1_024..16_384).step_by(2_048) {
905 reference[index] = 0.95;
906 if index + 1 < reference.len() {
907 reference[index + 1] = -0.7;
908 }
909 }
910 let test = vec![0.0; reference.len()];
911 let report = ArtifactReport::compare(&mono(reference), &mono(test)).unwrap();
912 assert!(
913 report.transient_loss_score > 0.2,
914 "transient-loss score: {}",
915 report.transient_loss_score
916 );
917 }
918
919 #[test]
920 fn detects_narrowband_musical_noise() {
921 let reference = vec![0.0; 16_384];
922 let test = (0..16_384)
923 .map(|index| (2.0 * PI * 1_000.0 * index as f64 / 16_000.0).sin() * 0.5)
924 .collect::<Vec<_>>();
925 let report = ArtifactReport::compare(&mono(reference), &mono(test)).unwrap();
926 assert!(
927 report.musical_noise_score > 0.1,
928 "musical-noise score: {}",
929 report.musical_noise_score
930 );
931 }
932
933 #[test]
934 fn detects_stereo_phase_inversion() {
935 let left = (0..16_384)
936 .map(|index| (2.0 * PI * 440.0 * index as f64 / 16_000.0).sin() * 0.5)
937 .collect::<Vec<_>>();
938 let right = left.clone();
939 let inverted = right.iter().map(|sample| -*sample).collect::<Vec<_>>();
940 let report =
941 ArtifactReport::compare(&stereo(left.clone(), right), &stereo(left, inverted)).unwrap();
942 assert!(
943 report.phase_distortion_score.unwrap_or(0.0) > 0.8,
944 "phase-distortion score: {:?}",
945 report.phase_distortion_score
946 );
947 }
948}