98 lines
3.5 KiB
Python
98 lines
3.5 KiB
Python
from __future__ import annotations
|
|
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from owner_voice_pet.audio_preprocess import SherpaOnnxDenoiserPreprocessor
|
|
from owner_voice_pet.models import AudioFrame, ErrorCode, ProviderError
|
|
from owner_voice_pet.speech_models import DEFAULT_DENOISER_PATH
|
|
|
|
|
|
class FakeDenoisedAudio:
|
|
def __init__(self, samples, sample_rate: int) -> None:
|
|
self.samples = samples
|
|
self.sample_rate = sample_rate
|
|
|
|
|
|
class FakeSherpa:
|
|
class OfflineSpeechDenoiserGtcrnModelConfig:
|
|
def __init__(self, model: str) -> None:
|
|
self.model = model
|
|
|
|
class OfflineSpeechDenoiserModelConfig:
|
|
def __init__(self, gtcrn, num_threads: int, provider: str) -> None:
|
|
self.gtcrn = gtcrn
|
|
self.num_threads = num_threads
|
|
self.provider = provider
|
|
|
|
class OnlineSpeechDenoiserConfig:
|
|
def __init__(self, model) -> None:
|
|
self.model = model
|
|
|
|
class OnlineSpeechDenoiser:
|
|
def __init__(self, config) -> None:
|
|
self.config = config
|
|
self.reset_called = False
|
|
|
|
def run(self, samples, sample_rate: int):
|
|
import numpy as np
|
|
|
|
return FakeDenoisedAudio(np.zeros_like(samples, dtype=np.float32), sample_rate)
|
|
|
|
def reset(self) -> None:
|
|
self.reset_called = True
|
|
|
|
|
|
class FailingSherpa(FakeSherpa):
|
|
class OnlineSpeechDenoiser(FakeSherpa.OnlineSpeechDenoiser):
|
|
def run(self, samples, sample_rate: int):
|
|
raise RuntimeError("denoise failed")
|
|
|
|
|
|
class AudioPreprocessTests(unittest.TestCase):
|
|
def test_denoiser_missing_model_is_structured(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FakeSherpa)
|
|
with self.assertRaises(ProviderError) as raised:
|
|
provider.load()
|
|
self.assertEqual(raised.exception.code, ErrorCode.NOISE_FILTER_MODEL_MISSING)
|
|
|
|
def test_denoiser_processes_frame_and_marks_metadata(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
model = Path(tmp) / DEFAULT_DENOISER_PATH
|
|
model.parent.mkdir(parents=True)
|
|
model.write_bytes(b"fake")
|
|
provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FakeSherpa)
|
|
provider.load()
|
|
frame = AudioFrame(
|
|
b"\xff\x7f",
|
|
16000,
|
|
1,
|
|
100,
|
|
7,
|
|
{"duration_ms": 20, "speech": True},
|
|
)
|
|
processed = provider.process_frame(frame)
|
|
|
|
self.assertNotEqual(processed.pcm, frame.pcm)
|
|
self.assertTrue(processed.metadata["denoised"])
|
|
self.assertEqual(processed.metadata["noise_filter_provider"], "sherpa_onnx_gtcrn")
|
|
self.assertEqual(processed.timestamp_ms, 100)
|
|
self.assertEqual(processed.frame_id, 7)
|
|
|
|
def test_denoiser_runtime_failure_is_structured(self) -> None:
|
|
with tempfile.TemporaryDirectory() as tmp:
|
|
model = Path(tmp) / DEFAULT_DENOISER_PATH
|
|
model.parent.mkdir(parents=True)
|
|
model.write_bytes(b"fake")
|
|
provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FailingSherpa)
|
|
provider.load()
|
|
with self.assertRaises(ProviderError) as raised:
|
|
provider.process_frame(AudioFrame(b"\xff\x7f", 16000, 1, 0, 0))
|
|
self.assertEqual(raised.exception.code, ErrorCode.NOISE_FILTER_FAILED)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|