Skip to main content

denoize/
benchmark.rs

1//! Objective quality and stereo-imaging reports.
2
3use crate::fft::{Complex, Fft};
4use crate::Audio;
5use std::f64::consts::PI;
6
7/// Normalized artifact-screening indicators derived from a reference/test pair.
8///
9/// Every score is in `[0, 1]`, where zero means that the corresponding
10/// artifact was not detected and one means a strong indication. These are
11/// deterministic, dependency-free screening signals rather than perceptual
12/// listening-test scores; use them to catch regressions and inspect the audio
13/// when a score rises.
14#[derive(Clone, Debug, PartialEq)]
15pub struct ArtifactReport {
16    /// Narrow spectral excesses in the test signal that are absent from the
17    /// reference (a common musical-noise / "birdie" signature).
18    pub musical_noise_score: f64,
19    /// Frame-to-frame modulation of the test/reference level ratio.
20    pub pumping_score: f64,
21    /// Fraction of strong reference onsets whose spectral flux is missing in
22    /// the test signal.
23    pub transient_loss_score: f64,
24    /// Inter-channel phase relationship error. `None` for mono input.
25    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    /// Dependency-free artifact-screening indicators.
40    pub artifact_scores: ArtifactReport,
41    /// Native STOI score in `[-1, 1]` when the input is long enough.
42    pub stoi: Option<f64>,
43    /// PESQ is `None` unless a separately licensed external adapter is used.
44    pub pesq: Option<f64>,
45    /// ViSQOL MOS-LQO (`[1, 5]`) when built with the `visqol` feature.
46    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    /// Compare a reference signal with a test signal and calculate artifact
106    /// screening indicators.
107    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    /// Return the machine-readable representation used by benchmark reports.
139    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            // Keep the HTML export in lockstep with every metric row in the
211            // Markdown report.  A fixed row count silently dropped the
212            // transient-loss and phase-distortion metrics when they were
213            // added to the comparison report.
214            .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
332/// Collect all four artifact signals in one shared STFT pass. The frame size
333/// is intentionally bounded so a long benchmark stays predictable in both
334/// memory and runtime.
335fn 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        // Musical noise is represented by narrow-band energy that is both
393        // absent from the reference and much stronger than its neighbours.
394        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    // At least 64 samples keeps the FFT meaningful for very short fixtures;
436    // longer inputs use a power-of-two frame no larger than 1024 samples.
437    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    // left * conjugate(right), preserving the inter-channel phase relation.
529    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}