| import io |
| import os |
| import re |
| import time |
| import wave |
|
|
| import modal |
| from pydantic import BaseModel |
|
|
|
|
| APP_NAME = "virtual-characters-tts" |
| GPU = os.environ.get("VC_TTS_GPU", "A10G") |
| BACKEND = os.environ.get("VC_TTS_BACKEND", "chatterbox") |
| LANGUAGE_ID = os.environ.get("VC_TTS_LANGUAGE_ID", "zh") |
| VOICE_DIR = "/voices" |
| HF_SECRET_NAME = os.environ.get("VC_HF_SECRET_NAME", "hf-token") |
| HF_SECRETS = [] if os.environ.get("VC_SKIP_HF_SECRET") == "1" else [modal.Secret.from_name(HF_SECRET_NAME)] |
|
|
|
|
| image = ( |
| modal.Image.debian_slim(python_version="3.11") |
| .apt_install("ffmpeg") |
| .uv_pip_install( |
| "chatterbox-tts>=0.1.6", |
| "fastapi[standard]>=0.115.0", |
| "kokoro>=0.9.2", |
| "ordered-set>=4.1.0", |
| "peft>=0.17.0", |
| "pypinyin>=0.55.0", |
| "soundfile>=0.12.1", |
| "torchaudio==2.6.0", |
| ) |
| .env({"VC_TTS_BACKEND": BACKEND, "VC_TTS_LANGUAGE_ID": LANGUAGE_ID}) |
| ) |
|
|
| voice_volume = modal.Volume.from_name("vc-tts-voices", create_if_missing=True) |
| app = modal.App(APP_NAME, image=image) |
|
|
|
|
| class TTSPayload(BaseModel): |
| text: str = "" |
| voice_id: str = "default" |
| emotion: str = "neutral" |
| speed: float = 1.0 |
| energy: float = 0.5 |
| audio_prompt_path: str | None = None |
| probe_only: bool = False |
|
|
|
|
| @app.cls( |
| gpu=GPU, |
| scaledown_window=60 * 3, |
| timeout=60 * 10, |
| secrets=HF_SECRETS, |
| volumes={VOICE_DIR: voice_volume}, |
| ) |
| class CharacterTTS: |
| def _ensure_loaded(self): |
| if getattr(self, "model", None) is not None: |
| return |
| import torch |
|
|
| self.backend = BACKEND |
| self.torch = torch |
| if self.backend == "kokoro": |
| from kokoro import KPipeline |
|
|
| self.model = KPipeline(lang_code="z") |
| self.sample_rate = 24000 |
| else: |
| try: |
| from chatterbox.mtl_tts import ChatterboxMultilingualTTS |
|
|
| self.model = ChatterboxMultilingualTTS.from_pretrained(device="cuda") |
| except Exception: |
| from chatterbox.tts import ChatterboxTTS |
|
|
| self.model = ChatterboxTTS.from_pretrained(device="cuda") |
| self.sample_rate = getattr(self.model, "sr", 24000) |
|
|
| @modal.method() |
| def health(self) -> dict: |
| return {"ok": True, "backend": BACKEND, "gpu": GPU, "loaded": getattr(self, "model", None) is not None} |
|
|
| @modal.fastapi_endpoint(method="GET") |
| async def health_http(self): |
| return self.health.local() |
|
|
| @modal.method() |
| def synthesize(self, text: str, voice_id: str = "default", emotion: str = "neutral", speed: float = 1.0, energy: float = 0.5) -> bytes: |
| return self._synthesize_bytes(text=text, voice_id=voice_id, emotion=emotion, speed=speed, energy=energy) |
|
|
| @modal.method() |
| def benchmark(self, text: str, voice_id: str = "default", emotion: str = "neutral", speed: float = 1.0, energy: float = 0.5) -> dict: |
| started = time.perf_counter() |
| was_loaded = getattr(self, "model", None) is not None |
| audio = self._synthesize_bytes(text=text, voice_id=voice_id, emotion=emotion, speed=speed, energy=energy) |
| elapsed = time.perf_counter() - started |
| audio_duration_s = _wav_duration_s(audio) |
| return { |
| "backend": BACKEND, |
| "gpu": GPU, |
| "text_chars": len(text), |
| "audio_bytes": len(audio), |
| "remote_s": round(elapsed, 3), |
| "audio_duration_s": audio_duration_s, |
| "real_time_factor": round(elapsed / audio_duration_s, 3) if audio_duration_s else None, |
| "was_loaded": was_loaded, |
| "sample_rate": getattr(self, "sample_rate", None), |
| "audio": audio, |
| } |
|
|
| @modal.method() |
| def benchmark_repeated( |
| self, |
| text: str, |
| repeats: int = 2, |
| voice_id: str = "default", |
| emotion: str = "neutral", |
| speed: float = 1.0, |
| energy: float = 0.5, |
| ) -> dict: |
| runs = [] |
| audio = b"" |
| for index in range(repeats): |
| result = self.benchmark.local( |
| text=text, |
| voice_id=voice_id, |
| emotion=emotion, |
| speed=speed, |
| energy=energy, |
| ) |
| audio = result.pop("audio") |
| result["index"] = index + 1 |
| runs.append(result) |
| return {"runs": runs, "audio": audio} |
|
|
| @modal.method() |
| def benchmark_sentence_stream( |
| self, |
| text: str, |
| voice_id: str = "default", |
| emotion: str = "neutral", |
| speed: float = 1.0, |
| energy: float = 0.5, |
| ) -> dict: |
| started = time.perf_counter() |
| sentences = _split_sentences(text) |
| events = [] |
| audio = b"" |
| for index, sentence in enumerate(sentences): |
| chunk_started = time.perf_counter() |
| audio = self._synthesize_bytes( |
| text=sentence, |
| voice_id=voice_id, |
| emotion=emotion, |
| speed=speed, |
| energy=energy, |
| ) |
| elapsed = time.perf_counter() - chunk_started |
| audio_duration_s = _wav_duration_s(audio) |
| events.append( |
| { |
| "index": index + 1, |
| "text": sentence, |
| "audio_bytes": len(audio), |
| "chunk_s": round(elapsed, 3), |
| "since_start_s": round(time.perf_counter() - started, 3), |
| "audio_duration_s": audio_duration_s, |
| "real_time_factor": round(elapsed / audio_duration_s, 3) if audio_duration_s else None, |
| "was_loaded_after_chunk": getattr(self, "model", None) is not None, |
| } |
| ) |
| return { |
| "backend": BACKEND, |
| "gpu": GPU, |
| "sentences": len(sentences), |
| "first_audio_s": events[0]["since_start_s"] if events else None, |
| "total_s": round(time.perf_counter() - started, 3), |
| "events": events, |
| "audio": audio, |
| } |
|
|
| @modal.fastapi_endpoint(method="POST") |
| async def tts(self, payload: TTSPayload): |
| from fastapi.responses import JSONResponse, StreamingResponse |
|
|
| data = payload.model_dump() |
| if data["probe_only"] or not data["text"].strip(): |
| return JSONResponse(self.health.local()) |
| audio = self._synthesize_bytes( |
| text=data["text"], |
| voice_id=data["voice_id"], |
| emotion=data["emotion"], |
| speed=float(data["speed"]), |
| energy=float(data["energy"]), |
| audio_prompt_path=data["audio_prompt_path"], |
| ) |
| return StreamingResponse(io.BytesIO(audio), media_type="audio/wav") |
|
|
| def _synthesize_bytes( |
| self, |
| text: str, |
| voice_id: str, |
| emotion: str, |
| speed: float, |
| energy: float, |
| audio_prompt_path: str | None = None, |
| ) -> bytes: |
| self._ensure_loaded() |
| import soundfile as sf |
| import torchaudio as ta |
|
|
| if self.backend == "kokoro": |
| voice = "zf_xiaoxiao" if voice_id == "default" else voice_id |
| generator = self.model(text, voice=voice) |
| chunks = [] |
| for _, _, audio in generator: |
| chunks.append(audio) |
| import numpy as np |
|
|
| wav = np.concatenate(chunks) if chunks else np.zeros(1, dtype="float32") |
| buffer = io.BytesIO() |
| sf.write(buffer, wav, self.sample_rate, format="WAV") |
| buffer.seek(0) |
| return buffer.read() |
|
|
| prompt = text |
| if emotion in {"happy", "smile"}: |
| prompt = f"{text} [chuckle]" |
| elif emotion in {"worried", "concerned", "sad"}: |
| prompt = f"[sigh] {text}" |
| elif emotion in {"playful"}: |
| prompt = f"{text} [laugh]" |
| elif emotion in {"battle_focus", "firm"}: |
| prompt = f"{text}" |
|
|
| kwargs = { |
| "language_id": LANGUAGE_ID, |
| "exaggeration": max(0.25, min(1.0, energy + 0.2)), |
| } |
| resolved_prompt = _resolve_voice_prompt(audio_prompt_path or voice_id) |
| if resolved_prompt: |
| kwargs["audio_prompt_path"] = resolved_prompt |
| if hasattr(self.model, "generate"): |
| try: |
| wav = self.model.generate(prompt, **kwargs) |
| except TypeError: |
| try: |
| wav = self.model.generate(prompt, language_id=LANGUAGE_ID) |
| except TypeError: |
| wav = self.model.generate(prompt) |
| else: |
| raise RuntimeError("Unsupported TTS model object") |
|
|
| buffer = io.BytesIO() |
| if hasattr(wav, "detach"): |
| wav = wav.detach().cpu() |
| ta.save(buffer, wav, self.sample_rate, format="wav") |
| buffer.seek(0) |
| return buffer.read() |
|
|
|
|
| def _split_sentences(text: str) -> list[str]: |
| parts = [part.strip() for part in re.split(r"(?<=[。!?!?;;])\s*", text) if part.strip()] |
| return parts or [text.strip()] |
|
|
|
|
| def _wav_duration_s(audio: bytes) -> float | None: |
| try: |
| with wave.open(io.BytesIO(audio), "rb") as wav_file: |
| frames = wav_file.getnframes() |
| rate = wav_file.getframerate() |
| return round(frames / float(rate), 3) if rate else None |
| except Exception: |
| return None |
|
|
|
|
| def _resolve_voice_prompt(voice_ref: str | None) -> str | None: |
| if not voice_ref: |
| return None |
| candidates = [ |
| os.path.join(VOICE_DIR, voice_ref), |
| os.path.join(VOICE_DIR, "prompts", voice_ref), |
| os.path.join(VOICE_DIR, f"{voice_ref}.wav"), |
| os.path.join(VOICE_DIR, "prompts", f"{voice_ref}.wav"), |
| ] |
| for candidate in candidates: |
| if os.path.exists(candidate): |
| return candidate |
| return None |
|
|
|
|
| @app.local_entrypoint() |
| def main( |
| text: str = "你好,我在听。", |
| output_path: str = "modal_tts_check.wav", |
| emotion: str = "neutral", |
| repeats: int = 1, |
| sentence_stream: bool = False, |
| ): |
| print(CharacterTTS().health.remote()) |
| started = time.perf_counter() |
| if sentence_stream: |
| result = CharacterTTS().benchmark_sentence_stream.remote(text=text, emotion=emotion) |
| elif repeats > 1: |
| result = CharacterTTS().benchmark_repeated.remote(text=text, emotion=emotion, repeats=repeats) |
| else: |
| result = CharacterTTS().benchmark.remote(text=text, emotion=emotion) |
| client_s = time.perf_counter() - started |
| audio = result.pop("audio") |
| with open(output_path, "wb") as f: |
| f.write(audio) |
| result["client_s"] = round(client_s, 3) |
| result["output_path"] = output_path |
| print(result) |
|
|