From 615aeaeb359ecc5e3e7a3e944fc53cdf85598162 Mon Sep 17 00:00:00 2001 From: devangkantharia Date: Mon, 10 Aug 2026 02:08:01 +0530 Subject: [PATCH] fix(kokoro): trim trailing silence and run-on noise on short prompt synthesis (#960) --- backend/backends/__init__.py | 1 + backend/backends/kokoro_backend.py | 7 +++- backend/tests/test_kokoro_trim.py | 66 ++++++++++++++++++++++++++++++ 3 files changed, 72 insertions(+), 2 deletions(-) create mode 100644 backend/tests/test_kokoro_trim.py diff --git a/backend/backends/__init__.py b/backend/backends/__init__.py index 142c430c..8f08c03f 100644 --- a/backend/backends/__init__.py +++ b/backend/backends/__init__.py @@ -369,6 +369,7 @@ def _get_non_qwen_tts_configs() -> list[ModelConfig]: engine="kokoro", hf_repo_id="hexgrad/Kokoro-82M", size_mb=350, + needs_trim=True, languages=["en", "es", "fr", "hi", "it", "pt", "ja", "zh"], ), ] diff --git a/backend/backends/kokoro_backend.py b/backend/backends/kokoro_backend.py index 0f6ff39e..4481cad7 100644 --- a/backend/backends/kokoro_backend.py +++ b/backend/backends/kokoro_backend.py @@ -288,7 +288,10 @@ class KokoroTTSBackend: # Return 1 second of silence as fallback return np.zeros(KOKORO_SAMPLE_RATE, dtype=np.float32), KOKORO_SAMPLE_RATE - audio = np.concatenate(audio_chunks) - return audio.astype(np.float32), KOKORO_SAMPLE_RATE + audio = np.concatenate(audio_chunks).astype(np.float32) + from ..utils.audio import trim_tts_output + + audio = trim_tts_output(audio, sample_rate=KOKORO_SAMPLE_RATE) + return audio, KOKORO_SAMPLE_RATE return await asyncio.to_thread(_generate_sync) diff --git a/backend/tests/test_kokoro_trim.py b/backend/tests/test_kokoro_trim.py new file mode 100644 index 00000000..550b561d --- /dev/null +++ b/backend/tests/test_kokoro_trim.py @@ -0,0 +1,66 @@ +"""Test Kokoro short prompt audio trimming and engine config.""" + +import numpy as np +import pytest + +from backend.backends import engine_needs_trim, get_model_config +from backend.utils.audio import trim_tts_output + + +def test_kokoro_engine_needs_trim_enabled(): + """Verify Kokoro engine is registered with needs_trim=True in model config.""" + assert engine_needs_trim("kokoro") is True + config = get_model_config("kokoro") + assert config is not None + assert config.needs_trim is True + + +def test_kokoro_trim_tts_output_removes_trailing_dead_space(): + """Verify trim_tts_output removes trailing silence past speech.""" + sr = 24000 + speech = np.full(int(sr * 1.5), 0.2, dtype=np.float32) # 1.5s speech + trailing_silence = np.zeros(int(sr * 1.0), dtype=np.float32) # 1.0s trailing dead space + raw_audio = np.concatenate([speech, trailing_silence]) + + trimmed = trim_tts_output(raw_audio, sample_rate=sr) + + # Trimming cuts trailing silence from 2.5s down to speech duration boundary (1.5s) + expected_dur_samples = int(sr * 1.5) + assert len(trimmed) == expected_dur_samples + assert len(trimmed) < len(raw_audio) + + +@pytest.mark.asyncio +async def test_kokoro_backend_generate_applies_trimming(monkeypatch): + """Verify KokoroTTSBackend.generate applies trimming on synthesized output.""" + from backend.backends.kokoro_backend import KokoroTTSBackend, KOKORO_SAMPLE_RATE + + backend = KokoroTTSBackend() + + # Mock _load_model_sync to avoid requiring real model load in pure unit test + monkeypatch.setattr(backend, "_load_model_sync", lambda: None) + monkeypatch.setattr(backend, "_model", object()) + + # Mock KPipeline output to yield audio with 1s trailing silence + sr = KOKORO_SAMPLE_RATE + speech = np.full(int(sr * 1.0), 0.2, dtype=np.float32) + silence = np.zeros(int(sr * 1.0), dtype=np.float32) + fake_audio = np.concatenate([speech, silence]) + + class FakeResult: + def __init__(self, audio): + self.audio = audio + + class FakePipeline: + def __call__(self, text, voice, speed=1.0): + yield FakeResult(fake_audio) + + monkeypatch.setattr(backend, "_get_pipeline", lambda lang: FakePipeline()) + + audio, sample_rate = await backend.generate("Read it back to me.", voice_prompt={}) + + assert sample_rate == sr + # Original fake audio was 2.0s (1s speech + 1s silence). + # Trimmed cuts trailing silence down to 1.0s speech boundary. + assert len(audio) / sr == pytest.approx(1.0, abs=0.05) + assert len(audio) < len(fake_audio)