voirs-evaluation 0.1.0-rc.1

Quality evaluation and assessment framework for VoiRS
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
//! SI-SDR (Scale-Invariant Signal-to-Distortion Ratio) Implementation
//!
//! Implementation of SI-SDR metric for evaluating speech enhancement and separation quality.
//! SI-SDR is particularly useful for measuring the quality of speech separation systems
//! and noise reduction algorithms.

use crate::EvaluationError;
use scirs2_core::ndarray::{Array1, Array2};
use std::collections::HashMap;
use voirs_sdk::{AudioBuffer, LanguageCode};

/// SI-SDR evaluation results with detailed breakdown
#[derive(Debug, Clone)]
pub struct SISdrResult {
    /// SI-SDR score in dB
    pub si_sdr_db: f32,
    /// SDR score in dB (for comparison)
    pub sdr_db: f32,
    /// Signal-to-Interference Ratio in dB
    pub sir_db: f32,
    /// Signal-to-Artifacts Ratio in dB
    pub sar_db: f32,
    /// Scale factor applied for SI-SDR calculation
    pub scale_factor: f32,
    /// Energy of the target signal
    pub target_energy: f32,
    /// Energy of the interference + artifacts
    pub distortion_energy: f32,
}

/// Batch SI-SDR evaluation results
#[derive(Debug, Clone)]
pub struct BatchSISdrResults {
    /// Individual results for each sample
    pub individual_results: Vec<SISdrResult>,
    /// Mean SI-SDR across all samples
    pub mean_si_sdr: f32,
    /// Standard deviation of SI-SDR scores
    pub std_si_sdr: f32,
    /// Median SI-SDR score
    pub median_si_sdr: f32,
    /// 95th percentile SI-SDR score
    pub percentile_95_si_sdr: f32,
    /// 5th percentile SI-SDR score
    pub percentile_5_si_sdr: f32,
    /// Number of samples processed
    pub num_samples: usize,
}

/// Language-specific SI-SDR configuration
#[derive(Debug, Clone)]
pub struct LanguageSISdrConfig {
    /// Language code
    pub language: LanguageCode,
    /// Frequency weighting for this language
    pub frequency_weights: Vec<f32>,
    /// SI-SDR threshold for acceptable quality
    pub quality_threshold_db: f32,
    /// Normalization factor for cross-language comparison
    pub normalization_factor: f32,
}

/// SI-SDR evaluator for speech quality assessment
pub struct SISdrEvaluator {
    /// Sample rate for processing
    sample_rate: u32,
    /// Whether to use zero-mean normalization
    zero_mean: bool,
    /// Language-specific configuration
    language_config: Option<LanguageSISdrConfig>,
}

impl SISdrEvaluator {
    /// Create new SI-SDR evaluator
    pub fn new(sample_rate: u32) -> Self {
        Self {
            sample_rate,
            zero_mean: true,
            language_config: None,
        }
    }

    /// Create SI-SDR evaluator with zero-mean option
    pub fn new_with_options(sample_rate: u32, zero_mean: bool) -> Self {
        Self {
            sample_rate,
            zero_mean,
            language_config: None,
        }
    }

    /// Set language-specific configuration
    pub fn set_language_config(&mut self, language: LanguageCode) {
        self.language_config = Some(Self::create_language_config(language));
    }

    /// Calculate SI-SDR between reference and estimated signals
    pub async fn calculate_si_sdr(
        &self,
        reference: &AudioBuffer,
        estimated: &AudioBuffer,
    ) -> Result<SISdrResult, EvaluationError> {
        // Validate inputs
        self.validate_inputs(reference, estimated)?;

        // Convert to arrays and ensure same length
        let min_len = reference.samples().len().min(estimated.samples().len());
        let ref_signal = Array1::from_vec(reference.samples()[..min_len].to_vec());
        let est_signal = Array1::from_vec(estimated.samples()[..min_len].to_vec());

        // Apply zero-mean normalization if enabled
        let (ref_normalized, est_normalized) = if self.zero_mean {
            let ref_mean = ref_signal.mean().unwrap_or(0.0);
            let est_mean = est_signal.mean().unwrap_or(0.0);
            (
                ref_signal.mapv(|x| x - ref_mean),
                est_signal.mapv(|x| x - est_mean),
            )
        } else {
            (ref_signal, est_signal)
        };

        // Calculate SI-SDR
        let si_sdr_result = self.compute_si_sdr(&ref_normalized, &est_normalized)?;

        // Calculate traditional SDR for comparison
        let sdr_result = self.compute_traditional_sdr(&ref_normalized, &est_normalized)?;

        // Calculate SIR and SAR
        let sir_result = self.compute_sir(&ref_normalized, &est_normalized)?;
        let sar_result = self.compute_sar(&ref_normalized, &est_normalized)?;

        Ok(SISdrResult {
            si_sdr_db: si_sdr_result.0,
            sdr_db: sdr_result,
            sir_db: sir_result,
            sar_db: sar_result,
            scale_factor: si_sdr_result.1,
            target_energy: si_sdr_result.2,
            distortion_energy: si_sdr_result.3,
        })
    }

    /// Calculate language-adapted SI-SDR with language-specific considerations
    pub async fn calculate_language_adapted_si_sdr(
        &self,
        reference: &AudioBuffer,
        estimated: &AudioBuffer,
        language: Option<LanguageCode>,
    ) -> Result<SISdrResult, EvaluationError> {
        // Use provided language or default from configuration
        let lang_config = if let Some(lang) = language {
            Self::create_language_config(lang)
        } else {
            self.language_config
                .clone()
                .unwrap_or_else(|| Self::create_language_config(LanguageCode::EnUs))
        };

        // Calculate base SI-SDR
        let mut result = self.calculate_si_sdr(reference, estimated).await?;

        // Apply language-specific calibration
        result.si_sdr_db = self.apply_language_calibration(result.si_sdr_db, &lang_config);
        result.sdr_db = self.apply_language_calibration(result.sdr_db, &lang_config);

        Ok(result)
    }

    /// Calculate SI-SDR for multiple signal pairs (batch processing)
    pub async fn calculate_batch_si_sdr(
        &self,
        reference_signals: &[AudioBuffer],
        estimated_signals: &[AudioBuffer],
    ) -> Result<BatchSISdrResults, EvaluationError> {
        if reference_signals.len() != estimated_signals.len() {
            return Err(EvaluationError::InvalidInput {
                message: "Number of reference and estimated signals must match".to_string(),
            });
        }

        if reference_signals.is_empty() {
            return Err(EvaluationError::InvalidInput {
                message: "At least one signal pair is required".to_string(),
            });
        }

        let mut individual_results = Vec::with_capacity(reference_signals.len());
        let mut si_sdr_scores = Vec::with_capacity(reference_signals.len());

        // Process each signal pair
        for (ref_signal, est_signal) in reference_signals.iter().zip(estimated_signals.iter()) {
            let result = self.calculate_si_sdr(ref_signal, est_signal).await?;
            si_sdr_scores.push(result.si_sdr_db);
            individual_results.push(result);
        }

        // Calculate statistics
        let mean_si_sdr = si_sdr_scores.iter().sum::<f32>() / si_sdr_scores.len() as f32;

        let variance = si_sdr_scores
            .iter()
            .map(|&score| (score - mean_si_sdr).powi(2))
            .sum::<f32>()
            / si_sdr_scores.len() as f32;
        let std_si_sdr = variance.sqrt();

        // Calculate percentiles
        let mut sorted_scores = si_sdr_scores.clone();
        sorted_scores.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));

        let median_si_sdr = if sorted_scores.len() % 2 == 0 {
            let mid = sorted_scores.len() / 2;
            (sorted_scores[mid - 1] + sorted_scores[mid]) / 2.0
        } else {
            sorted_scores[sorted_scores.len() / 2]
        };

        let percentile_95_idx =
            ((sorted_scores.len() as f32 * 0.95) as usize).min(sorted_scores.len() - 1);
        let percentile_5_idx = (sorted_scores.len() as f32 * 0.05) as usize;

        let percentile_95_si_sdr = sorted_scores[percentile_95_idx];
        let percentile_5_si_sdr = sorted_scores[percentile_5_idx];

        Ok(BatchSISdrResults {
            individual_results,
            mean_si_sdr,
            std_si_sdr,
            median_si_sdr,
            percentile_95_si_sdr,
            percentile_5_si_sdr,
            num_samples: reference_signals.len(),
        })
    }

    /// Validate input audio buffers
    fn validate_inputs(
        &self,
        reference: &AudioBuffer,
        estimated: &AudioBuffer,
    ) -> Result<(), EvaluationError> {
        if reference.sample_rate() != self.sample_rate {
            return Err(EvaluationError::InvalidInput {
                message: format!(
                    "Reference signal sample rate {} doesn't match evaluator rate {}",
                    reference.sample_rate(),
                    self.sample_rate
                ),
            });
        }

        if estimated.sample_rate() != self.sample_rate {
            return Err(EvaluationError::InvalidInput {
                message: format!(
                    "Estimated signal sample rate {} doesn't match evaluator rate {}",
                    estimated.sample_rate(),
                    self.sample_rate
                ),
            });
        }

        if reference.channels() != 1 || estimated.channels() != 1 {
            return Err(EvaluationError::InvalidInput {
                message: "SI-SDR requires mono audio".to_string(),
            });
        }

        if reference.samples().is_empty() || estimated.samples().is_empty() {
            return Err(EvaluationError::InvalidInput {
                message: "Audio signals cannot be empty".to_string(),
            });
        }

        Ok(())
    }

    /// Compute SI-SDR metric
    fn compute_si_sdr(
        &self,
        reference: &Array1<f32>,
        estimated: &Array1<f32>,
    ) -> Result<(f32, f32, f32, f32), EvaluationError> {
        if reference.len() != estimated.len() {
            return Err(EvaluationError::InvalidInput {
                message: "Reference and estimated signals must have the same length".to_string(),
            });
        }

        // Calculate optimal scaling factor (projection)
        let dot_product = reference.dot(estimated);
        let reference_power = reference.dot(reference);

        if reference_power < 1e-12f32 {
            return Err(EvaluationError::AudioProcessingError {
                message: "Reference signal has zero power".to_string(),
                source: None,
            });
        }

        let scale_factor = dot_product / reference_power;

        // Calculate scaled reference (target)
        let scaled_reference = reference.mapv(|x| x * scale_factor);

        // Calculate distortion (estimated - scaled_reference)
        let distortion = estimated - &scaled_reference;

        // Calculate energies
        let target_energy = scaled_reference.dot(&scaled_reference);
        let distortion_energy = distortion.dot(&distortion);

        // Calculate SI-SDR in dB
        let si_sdr_db = if distortion_energy > 1e-12f32 {
            10.0 * (target_energy / distortion_energy).log10()
        } else {
            // Very small distortion, return high SI-SDR
            100.0
        };

        Ok((si_sdr_db, scale_factor, target_energy, distortion_energy))
    }

    /// Compute traditional SDR metric for comparison
    fn compute_traditional_sdr(
        &self,
        reference: &Array1<f32>,
        estimated: &Array1<f32>,
    ) -> Result<f32, EvaluationError> {
        // Traditional SDR uses the original reference without scaling
        let distortion = estimated - reference;

        let signal_power = reference.dot(reference);
        let distortion_power = distortion.dot(&distortion);

        if distortion_power < 1e-12f32 {
            return Ok(100.0); // Very small distortion
        }

        if signal_power < 1e-12f32 {
            return Err(EvaluationError::AudioProcessingError {
                message: "Reference signal has zero power".to_string(),
                source: None,
            });
        }

        let sdr_db = 10.0 * (signal_power / distortion_power).log10();
        Ok(sdr_db)
    }

    /// Compute Signal-to-Interference Ratio (SIR)
    fn compute_sir(
        &self,
        reference: &Array1<f32>,
        estimated: &Array1<f32>,
    ) -> Result<f32, EvaluationError> {
        // For single-source scenarios, SIR is similar to SI-SDR
        // In multi-source scenarios, this would measure interference from other sources
        let (si_sdr_db, _, _, _) = self.compute_si_sdr(reference, estimated)?;

        // For simplicity, return SI-SDR as SIR estimate
        // In practice, this would require knowledge of interference sources
        Ok(si_sdr_db)
    }

    /// Compute Signal-to-Artifacts Ratio (SAR)
    fn compute_sar(
        &self,
        reference: &Array1<f32>,
        estimated: &Array1<f32>,
    ) -> Result<f32, EvaluationError> {
        // Calculate artifacts as high-frequency distortion
        let distortion = estimated - reference;

        // Apply high-pass filtering to extract artifacts
        let artifacts = self.apply_highpass_filter(&distortion)?;

        let signal_power = reference.dot(reference);
        let artifacts_power = artifacts.dot(&artifacts);

        if artifacts_power < 1e-12f32 {
            return Ok(100.0); // Very low artifacts
        }

        if signal_power < 1e-12f32 {
            return Err(EvaluationError::AudioProcessingError {
                message: "Reference signal has zero power".to_string(),
                source: None,
            });
        }

        let sar_db = 10.0 * (signal_power / artifacts_power).log10();
        Ok(sar_db)
    }

    /// Apply simple high-pass filter for artifact detection
    fn apply_highpass_filter(&self, signal: &Array1<f32>) -> Result<Array1<f32>, EvaluationError> {
        if signal.len() < 2 {
            return Ok(signal.clone());
        }

        let mut filtered = Array1::zeros(signal.len());

        // Simple first-order high-pass filter: y[n] = x[n] - x[n-1]
        filtered[0] = signal[0];
        for i in 1..signal.len() {
            filtered[i] = signal[i] - signal[i - 1];
        }

        Ok(filtered)
    }

    /// Create language-specific configuration
    fn create_language_config(language: LanguageCode) -> LanguageSISdrConfig {
        match language {
            LanguageCode::EnUs | LanguageCode::EnGb => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.0; 20],
                quality_threshold_db: 10.0,
                normalization_factor: 1.0,
            },
            LanguageCode::JaJp => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.1; 20], // Slightly higher weight for Japanese
                quality_threshold_db: 9.5,
                normalization_factor: 0.98,
            },
            LanguageCode::ZhCn => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.05; 20], // Tonal language adjustment
                quality_threshold_db: 9.0,
                normalization_factor: 0.96,
            },
            LanguageCode::EsEs | LanguageCode::EsMx => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.02; 20],
                quality_threshold_db: 10.2,
                normalization_factor: 1.01,
            },
            LanguageCode::FrFr => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.03; 20],
                quality_threshold_db: 10.1,
                normalization_factor: 1.005,
            },
            LanguageCode::DeDe => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.01; 20],
                quality_threshold_db: 10.3,
                normalization_factor: 1.02,
            },
            _ => LanguageSISdrConfig {
                language,
                frequency_weights: vec![1.0; 20],
                quality_threshold_db: 10.0,
                normalization_factor: 1.0,
            },
        }
    }

    /// Apply language-specific calibration
    fn apply_language_calibration(&self, base_score: f32, config: &LanguageSISdrConfig) -> f32 {
        // Apply normalization factor
        let calibrated = base_score * config.normalization_factor;

        // Apply threshold-based adjustment
        if calibrated < config.quality_threshold_db {
            calibrated * 0.95 // Penalty for below-threshold quality
        } else {
            calibrated * 1.02 // Slight bonus for above-threshold quality
        }
    }

    /// Calculate improvement in SI-SDR (useful for enhancement systems)
    pub async fn calculate_si_sdr_improvement(
        &self,
        noisy: &AudioBuffer,
        enhanced: &AudioBuffer,
        clean: &AudioBuffer,
    ) -> Result<f32, EvaluationError> {
        // Calculate SI-SDR before enhancement (noisy vs clean)
        let before_result = self.calculate_si_sdr(clean, noisy).await?;

        // Calculate SI-SDR after enhancement (enhanced vs clean)
        let after_result = self.calculate_si_sdr(clean, enhanced).await?;

        // Return improvement
        Ok(after_result.si_sdr_db - before_result.si_sdr_db)
    }

    /// Check if SI-SDR score meets quality threshold for a given language
    pub fn meets_quality_threshold(&self, si_sdr_db: f32, language: Option<LanguageCode>) -> bool {
        let config = if let Some(lang) = language {
            Self::create_language_config(lang)
        } else {
            self.language_config
                .clone()
                .unwrap_or_else(|| Self::create_language_config(LanguageCode::EnUs))
        };

        si_sdr_db >= config.quality_threshold_db
    }

    /// Get sample rate
    pub fn sample_rate(&self) -> u32 {
        self.sample_rate
    }

    /// Check if zero-mean normalization is enabled
    pub fn zero_mean_enabled(&self) -> bool {
        self.zero_mean
    }
}

impl SISdrResult {
    /// Check if the result indicates good quality
    pub fn is_good_quality(&self, threshold_db: Option<f32>) -> bool {
        let threshold = threshold_db.unwrap_or(10.0);
        self.si_sdr_db >= threshold
    }

    /// Get quality category as string
    pub fn quality_category(&self) -> &'static str {
        match self.si_sdr_db {
            x if x >= 20.0 => "Excellent",
            x if x >= 15.0 => "Good",
            x if x >= 10.0 => "Fair",
            x if x >= 5.0 => "Poor",
            _ => "Very Poor",
        }
    }

    /// Format result as human-readable string
    pub fn format_result(&self) -> String {
        format!(
            "SI-SDR: {:.2} dB ({}), SDR: {:.2} dB, SIR: {:.2} dB, SAR: {:.2} dB",
            self.si_sdr_db,
            self.quality_category(),
            self.sdr_db,
            self.sir_db,
            self.sar_db
        )
    }
}

impl BatchSISdrResults {
    /// Get summary statistics as string
    pub fn summary(&self) -> String {
        format!(
            "Batch SI-SDR Results ({} samples):\n\
             Mean: {:.2} dB, Std: {:.2} dB, Median: {:.2} dB\n\
             95th percentile: {:.2} dB, 5th percentile: {:.2} dB",
            self.num_samples,
            self.mean_si_sdr,
            self.std_si_sdr,
            self.median_si_sdr,
            self.percentile_95_si_sdr,
            self.percentile_5_si_sdr
        )
    }

    /// Get percentage of samples above threshold
    pub fn percentage_above_threshold(&self, threshold_db: f32) -> f32 {
        let count_above = self
            .individual_results
            .iter()
            .filter(|result| result.si_sdr_db >= threshold_db)
            .count();

        (count_above as f32 / self.num_samples as f32) * 100.0
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::f32::consts::PI;
    use voirs_sdk::AudioBuffer;

    #[tokio::test]
    async fn test_si_sdr_evaluator_creation() {
        let evaluator = SISdrEvaluator::new(16000);
        assert_eq!(evaluator.sample_rate(), 16000);
        assert!(evaluator.zero_mean_enabled());

        let evaluator_no_zero_mean = SISdrEvaluator::new_with_options(16000, false);
        assert!(!evaluator_no_zero_mean.zero_mean_enabled());
    }

    #[tokio::test]
    async fn test_perfect_si_sdr() {
        let evaluator = SISdrEvaluator::new(16000);

        // Perfect reconstruction should give very high SI-SDR
        let samples = vec![0.1, 0.2, 0.3, 0.4, 0.5];
        let reference = AudioBuffer::new(samples.clone(), 16000, 1);
        let estimated = AudioBuffer::new(samples, 16000, 1);

        let result = evaluator
            .calculate_si_sdr(&reference, &estimated)
            .await
            .unwrap();

        // Perfect match should give very high SI-SDR
        assert!(result.si_sdr_db > 50.0);
        assert_eq!(result.quality_category(), "Excellent");
    }

    #[tokio::test]
    async fn test_scaled_signal_si_sdr() {
        let evaluator = SISdrEvaluator::new(16000);

        // Test with scaled version of the signal
        let reference_samples = vec![0.1, 0.2, 0.3, 0.4, 0.5];
        let estimated_samples: Vec<f32> = reference_samples.iter().map(|x| x * 2.0).collect();

        let reference = AudioBuffer::new(reference_samples, 16000, 1);
        let estimated = AudioBuffer::new(estimated_samples, 16000, 1);

        let result = evaluator
            .calculate_si_sdr(&reference, &estimated)
            .await
            .unwrap();

        // Scaled signal should still give very high SI-SDR
        assert!(result.si_sdr_db > 50.0);
        assert!((result.scale_factor - 2.0).abs() < 0.001);
    }

    #[tokio::test]
    async fn test_noisy_signal_si_sdr() {
        let evaluator = SISdrEvaluator::new(16000);

        // Create a clean signal
        let duration_samples = 1000;
        let mut clean_samples = Vec::with_capacity(duration_samples);
        for i in 0..duration_samples {
            let t = i as f32 / 16000.0;
            clean_samples.push(0.5 * (2.0 * PI * 440.0 * t).sin());
        }

        // Add noise to create estimated signal
        let mut noisy_samples = clean_samples.clone();
        for sample in &mut noisy_samples {
            *sample += 0.1 * (scirs2_core::random::random::<f32>() - 0.5);
        }

        let reference = AudioBuffer::new(clean_samples, 16000, 1);
        let estimated = AudioBuffer::new(noisy_samples, 16000, 1);

        let result = evaluator
            .calculate_si_sdr(&reference, &estimated)
            .await
            .unwrap();

        // Noisy signal should have lower SI-SDR but still positive
        assert!(result.si_sdr_db > 0.0);
        assert!(result.si_sdr_db < 30.0);
    }

    #[tokio::test]
    async fn test_si_sdr_improvement() {
        let evaluator = SISdrEvaluator::new(16000);

        // Create clean, noisy, and enhanced signals
        let clean_samples = vec![0.5, 0.3, 0.1, -0.1, -0.3, -0.5];
        let noisy_samples = vec![0.7, 0.5, 0.3, 0.1, -0.1, -0.3]; // More significant noise
        let enhanced_samples = vec![0.51, 0.29, 0.11, -0.09, -0.29, -0.51]; // Better reconstruction

        let clean = AudioBuffer::new(clean_samples, 16000, 1);
        let noisy = AudioBuffer::new(noisy_samples, 16000, 1);
        let enhanced = AudioBuffer::new(enhanced_samples, 16000, 1);

        let improvement = evaluator
            .calculate_si_sdr_improvement(&noisy, &enhanced, &clean)
            .await
            .unwrap();

        // Enhancement should provide improvement (can be positive or negative)
        // Just verify the calculation works correctly
        let before_result = evaluator.calculate_si_sdr(&clean, &noisy).await.unwrap();
        let after_result = evaluator.calculate_si_sdr(&clean, &enhanced).await.unwrap();
        let expected_improvement = after_result.si_sdr_db - before_result.si_sdr_db;

        assert!((improvement - expected_improvement).abs() < 1e-6);
    }

    #[tokio::test]
    async fn test_language_adapted_si_sdr() {
        let evaluator = SISdrEvaluator::new(16000);

        let reference_samples = vec![0.1, 0.2, 0.3, 0.4, 0.5];
        let estimated_samples = vec![0.12, 0.21, 0.29, 0.39, 0.48];

        let reference = AudioBuffer::new(reference_samples, 16000, 1);
        let estimated = AudioBuffer::new(estimated_samples, 16000, 1);

        // Test with different languages
        let en_result = evaluator
            .calculate_language_adapted_si_sdr(&reference, &estimated, Some(LanguageCode::EnUs))
            .await
            .unwrap();

        let ja_result = evaluator
            .calculate_language_adapted_si_sdr(&reference, &estimated, Some(LanguageCode::JaJp))
            .await
            .unwrap();

        // Both should be valid, but potentially different due to language calibration
        assert!(en_result.si_sdr_db > 0.0);
        assert!(ja_result.si_sdr_db > 0.0);
    }

    #[tokio::test]
    async fn test_batch_si_sdr() {
        let evaluator = SISdrEvaluator::new(16000);

        // Create multiple signal pairs
        let reference_signals = vec![
            AudioBuffer::new(vec![0.1, 0.2, 0.3], 16000, 1),
            AudioBuffer::new(vec![0.4, 0.5, 0.6], 16000, 1),
            AudioBuffer::new(vec![0.7, 0.8, 0.9], 16000, 1),
        ];

        let estimated_signals = vec![
            AudioBuffer::new(vec![0.11, 0.19, 0.31], 16000, 1),
            AudioBuffer::new(vec![0.39, 0.51, 0.59], 16000, 1),
            AudioBuffer::new(vec![0.71, 0.79, 0.89], 16000, 1),
        ];

        let batch_results = evaluator
            .calculate_batch_si_sdr(&reference_signals, &estimated_signals)
            .await
            .unwrap();

        assert_eq!(batch_results.num_samples, 3);
        assert_eq!(batch_results.individual_results.len(), 3);
        assert!(batch_results.mean_si_sdr > 0.0);
        assert!(batch_results.std_si_sdr >= 0.0);
    }

    #[test]
    fn test_quality_threshold() {
        let evaluator = SISdrEvaluator::new(16000);

        assert!(evaluator.meets_quality_threshold(15.0, Some(LanguageCode::EnUs)));
        assert!(!evaluator.meets_quality_threshold(5.0, Some(LanguageCode::EnUs)));
    }

    #[test]
    fn test_si_sdr_result_methods() {
        let result = SISdrResult {
            si_sdr_db: 15.5,
            sdr_db: 14.2,
            sir_db: 16.1,
            sar_db: 18.3,
            scale_factor: 1.1,
            target_energy: 0.5,
            distortion_energy: 0.05,
        };

        assert!(result.is_good_quality(Some(10.0)));
        assert_eq!(result.quality_category(), "Good");

        let formatted = result.format_result();
        assert!(formatted.contains("15.5"));
        assert!(formatted.contains("Good"));
    }

    #[test]
    fn test_batch_results_methods() {
        let individual_results = vec![
            SISdrResult {
                si_sdr_db: 15.0,
                sdr_db: 14.0,
                sir_db: 16.0,
                sar_db: 18.0,
                scale_factor: 1.0,
                target_energy: 0.5,
                distortion_energy: 0.05,
            },
            SISdrResult {
                si_sdr_db: 8.0,
                sdr_db: 7.5,
                sir_db: 8.5,
                sar_db: 10.0,
                scale_factor: 1.1,
                target_energy: 0.4,
                distortion_energy: 0.08,
            },
        ];

        let batch_results = BatchSISdrResults {
            individual_results,
            mean_si_sdr: 11.5,
            std_si_sdr: 3.5,
            median_si_sdr: 11.5,
            percentile_95_si_sdr: 15.0,
            percentile_5_si_sdr: 8.0,
            num_samples: 2,
        };

        let summary = batch_results.summary();
        assert!(summary.contains("11.5"));
        assert!(summary.contains("2 samples"));

        let percentage_above = batch_results.percentage_above_threshold(10.0);
        assert_eq!(percentage_above, 50.0); // 1 out of 2 samples above 10.0 dB
    }

    #[test]
    fn test_language_config_creation() {
        let en_config = SISdrEvaluator::create_language_config(LanguageCode::EnUs);
        let ja_config = SISdrEvaluator::create_language_config(LanguageCode::JaJp);

        assert_eq!(en_config.language, LanguageCode::EnUs);
        assert_eq!(ja_config.language, LanguageCode::JaJp);

        // Different languages should have different thresholds
        assert_ne!(
            en_config.quality_threshold_db,
            ja_config.quality_threshold_db
        );
    }
}