Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 79 additions & 20 deletions src/voxcpm/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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.

Expand All @@ -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``,
Expand Down Expand Up @@ -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:
Expand Down
189 changes: 189 additions & 0 deletions tests/test_long_text_chunking.py
Original file line number Diff line number Diff line change
@@ -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"]))