diff --git a/src/voxcpm/core.py b/src/voxcpm/core.py index 1a1d8398..9b00d5c1 100644 --- a/src/voxcpm/core.py +++ b/src/voxcpm/core.py @@ -11,6 +11,51 @@ from .model.utils import next_and_close +# Autoregressive quality degrades on long inputs: past a few hundred characters +# the output drifts into distortion/instability (issue #372, worse on VoxCPM2). +# Splitting the text into sentence-sized chunks and synthesizing each with the +# same prompt cache keeps every generation short while preserving the voice. +DEFAULT_MAX_CHUNK_CHARS = 200 +# Sentence-ending marks used as preferred split points (CJK + ASCII). +_SENTENCE_DELIMITERS = "。!?;…!?;" + + +def split_text_into_chunks(text: str, max_chars: int = DEFAULT_MAX_CHUNK_CHARS) -> list: + """Split ``text`` into chunks no longer than ``max_chars`` characters. + + Splits preferentially after sentence-ending punctuation so chunk boundaries + fall on natural pauses; a single sentence longer than ``max_chars`` is + hard-sliced. Returns ``[text]`` unchanged when it already fits (or when + ``max_chars`` is non-positive, which disables splitting). + """ + text = text.strip() + if not text: + return [] + if max_chars <= 0 or len(text) <= max_chars: + return [text] + + # Zero-width split after each run of sentence-ending marks, keeping them. + pieces = re.split(f"(?<=[{re.escape(_SENTENCE_DELIMITERS)}])", text) + chunks = [] + buf = "" + for piece in pieces: + if not piece: + continue + if len(buf) + len(piece) <= max_chars: + buf += piece + continue + if buf: + chunks.append(buf.strip()) + buf = "" + while len(piece) > max_chars: + chunks.append(piece[:max_chars].strip()) + piece = piece[max_chars:] + buf = piece + if buf.strip(): + chunks.append(buf.strip()) + return [chunk for chunk in chunks if chunk] + + class VoxCPM: def __init__( self, @@ -197,6 +242,7 @@ def _generate( retry_badcase_ratio_threshold: float = 6.0, streaming: bool = False, seed: Optional[int] = None, + max_chunk_chars: int = DEFAULT_MAX_CHUNK_CHARS, ) -> Generator[np.ndarray, None, None]: """Synthesize speech for the given text and return a single waveform. @@ -220,6 +266,11 @@ def _generate( retry_badcase_ratio_threshold: Threshold for audio-to-text ratio. streaming: Whether to return a generator of audio chunks. seed: Optional random seed for reproducibility. + max_chunk_chars: Long inputs are split into chunks of at most this + many characters (on sentence boundaries) and synthesized + sequentially with the same prompt, avoiding the quality + degradation seen on long single-shot generations. Set to 0 or a + negative value to disable splitting. Returns: Generator of numpy.ndarray: 1D waveform array (float32) on CPU. Yields audio chunks for each generation step if ``streaming=True``, @@ -285,29 +336,37 @@ def _generate( self.text_normalizer = TextNormalizer() text = self.text_normalizer.normalize(text) - generate_result = self.tts_model._generate_with_prompt_cache( - target_text=text, - prompt_cache=fixed_prompt_cache, - min_len=min_len, - max_len=max_len, - inference_timesteps=inference_timesteps, - cfg_value=cfg_value, - retry_badcase=retry_badcase, - retry_badcase_max_times=retry_badcase_max_times, - retry_badcase_ratio_threshold=retry_badcase_ratio_threshold, - streaming=streaming, - seed=seed, - ) + chunks = split_text_into_chunks(text, max_chars=max_chunk_chars) or [text] + + def _synthesize(chunk_text: str): + return self.tts_model._generate_with_prompt_cache( + target_text=chunk_text, + prompt_cache=fixed_prompt_cache, + min_len=min_len, + max_len=max_len, + inference_timesteps=inference_timesteps, + cfg_value=cfg_value, + retry_badcase=retry_badcase, + retry_badcase_max_times=retry_badcase_max_times, + retry_badcase_ratio_threshold=retry_badcase_ratio_threshold, + streaming=streaming, + seed=seed, + ) if streaming: - try: - for wav, _, _ in generate_result: - yield wav.squeeze(0).cpu().numpy() - finally: - generate_result.close() + for chunk_text in chunks: + generate_result = _synthesize(chunk_text) + try: + for wav, _, _ in generate_result: + yield wav.squeeze(0).cpu().numpy() + finally: + generate_result.close() else: - wav, _, _ = next_and_close(generate_result) - yield wav.squeeze(0).cpu().numpy() + parts = [] + for chunk_text in chunks: + wav, _, _ = next_and_close(_synthesize(chunk_text)) + parts.append(wav.squeeze(0).cpu().numpy()) + yield parts[0] if len(parts) == 1 else np.concatenate(parts) finally: for tmp_path in temp_files: diff --git a/tests/test_long_text_chunking.py b/tests/test_long_text_chunking.py new file mode 100644 index 00000000..f118d3da --- /dev/null +++ b/tests/test_long_text_chunking.py @@ -0,0 +1,189 @@ +"""Regression tests for long-text chunking (issue #372). + +Long single-shot generations drift into distortion past a few hundred +characters. ``VoxCPM._generate`` now splits long text into sentence-sized +chunks and synthesizes each one separately, so every generation stays short. + +The heavy model dependencies are stubbed so ``core.py`` can be loaded in +isolation without torch / the model weights. +""" + +from __future__ import annotations + +import importlib.util +import sys +import types +from pathlib import Path + +import numpy as np +import pytest + +ROOT = Path(__file__).resolve().parents[1] +CORE_PATH = ROOT / "src" / "voxcpm" / "core.py" + +MAX_CHARS = 40 # small threshold so tests stay short and deterministic + + +def _load_core(): + """Load voxcpm.core with its heavy imports replaced by lightweight stubs.""" + pkg = types.ModuleType("voxcpm") + pkg.__path__ = [str(ROOT / "src" / "voxcpm")] + sys.modules["voxcpm"] = pkg + + hub = types.ModuleType("huggingface_hub") + hub.snapshot_download = lambda *a, **k: None + sys.modules["huggingface_hub"] = hub + + model_pkg = types.ModuleType("voxcpm.model") + model_pkg.__path__ = [] + sys.modules["voxcpm.model"] = model_pkg + + voxcpm_mod = types.ModuleType("voxcpm.model.voxcpm") + voxcpm_mod.VoxCPMModel = type("VoxCPMModel", (), {}) + voxcpm_mod.LoRAConfig = type("LoRAConfig", (), {}) + sys.modules["voxcpm.model.voxcpm"] = voxcpm_mod + + voxcpm2_mod = types.ModuleType("voxcpm.model.voxcpm2") + voxcpm2_mod.VoxCPM2Model = type("VoxCPM2Model", (), {}) + sys.modules["voxcpm.model.voxcpm2"] = voxcpm2_mod + + utils_mod = types.ModuleType("voxcpm.model.utils") + + def next_and_close(gen): + try: + return next(gen) + finally: + gen.close() + + utils_mod.next_and_close = next_and_close + sys.modules["voxcpm.model.utils"] = utils_mod + + spec = importlib.util.spec_from_file_location("voxcpm.core", CORE_PATH) + core = importlib.util.module_from_spec(spec) + sys.modules["voxcpm.core"] = core + assert spec.loader is not None + spec.loader.exec_module(core) + return core + + +core = _load_core() + + +class _FakeWav: + """Stand-in for a torch tensor: squeeze(0).cpu().numpy() -> np.ndarray.""" + + def __init__(self, arr): + self._arr = arr + + def squeeze(self, _dim): + return self + + def cpu(self): + return self + + def numpy(self): + return self._arr + + +class _RecordingModel: + """Records every target_text handed to the generation entry point.""" + + def __init__(self): + self.calls = [] + + def _generate_with_prompt_cache(self, target_text, streaming=False, **kwargs): + self.calls.append(target_text) + # one sample per character -> audio length reflects chunk length + wav = _FakeWav(np.zeros(len(target_text), dtype=np.float32)) + + def _gen(): + yield wav, None, None + + return _gen() + + +def _make_vox(model): + vox = core.VoxCPM.__new__(core.VoxCPM) + vox.tts_model = model + vox.denoiser = None + vox.text_normalizer = None + return vox + + +# --------------------------------------------------------------------------- # +# split_text_into_chunks (pure helper) +# --------------------------------------------------------------------------- # + + +def test_short_text_is_single_chunk(): + text = "你好世界。" + assert core.split_text_into_chunks(text, max_chars=MAX_CHARS) == [text] + + +def test_long_text_splits_on_sentence_boundaries(): + text = "第一句话。" * 20 # 100 chars, far over MAX_CHARS + chunks = core.split_text_into_chunks(text, max_chars=MAX_CHARS) + assert len(chunks) > 1, "long text must be split into multiple chunks" + assert all(len(c) <= MAX_CHARS for c in chunks), "no chunk may exceed max_chars" + assert "".join(chunks) == text, "chunks must reconstruct the original text" + + +def test_splitting_disabled_when_max_chars_non_positive(): + text = "很长的句子" * 50 + assert core.split_text_into_chunks(text, max_chars=0) == [text] + + +def test_sentence_without_delimiter_is_hard_sliced(): + text = "字" * 100 # no punctuation to split on + chunks = core.split_text_into_chunks(text, max_chars=MAX_CHARS) + assert all(len(c) <= MAX_CHARS for c in chunks) + assert "".join(chunks) == text + + +# --------------------------------------------------------------------------- # +# _generate wiring: long text -> multiple model calls, audio concatenated +# --------------------------------------------------------------------------- # + + +def test_generate_chunks_long_text_into_multiple_calls(): + model = _RecordingModel() + vox = _make_vox(model) + text = "这是一段很长的测试文本。" * 10 # ~120 chars + + out = list(vox._generate(text, streaming=False, max_chunk_chars=MAX_CHARS)) + + assert len(model.calls) > 1, ( + f"long text should produce >1 generation call, got {len(model.calls)}" + ) + assert all(len(c) <= MAX_CHARS for c in model.calls) + # single concatenated waveform whose length == total synthesized characters + assert len(out) == 1 + assert out[0].shape[0] == sum(len(c) for c in model.calls) + print(f"long text -> {len(model.calls)} chunks, audio samples={out[0].shape[0]}") + + +def test_generate_short_text_is_single_call(): + model = _RecordingModel() + vox = _make_vox(model) + text = "短文本。" + + list(vox._generate(text, streaming=False, max_chunk_chars=MAX_CHARS)) + + assert len(model.calls) == 1, "short text must not be split" + print(f"short text -> {len(model.calls)} chunk") + + +def test_generate_streaming_yields_each_chunk(): + model = _RecordingModel() + vox = _make_vox(model) + text = "流式测试的句子。" * 8 + + out = list(vox._generate(text, streaming=True, max_chunk_chars=MAX_CHARS)) + + assert len(model.calls) > 1 + assert len(out) == len(model.calls), "streaming yields one array per chunk" + print(f"streaming long text -> {len(model.calls)} chunks yielded") + + +if __name__ == "__main__": + sys.exit(pytest.main([__file__, "-v", "-s"]))