CoolFace
Apppublic

Mike0021/zonos2

sourceHugging Faceupdated 4mo agoView on Hugging Face
3likes
app.py326 linesDownload Raw Back to root
1import os2import tempfile3import threading4import time5import wave6from pathlib import Path7 8os.environ.setdefault("HF_HOME", "/tmp/.cache/huggingface")9os.environ.setdefault("HF_MODULES_CACHE", "/tmp/hf_modules")10os.environ.setdefault("MPLCONFIGDIR", "/tmp/matplotlib")11os.environ.setdefault("ZONOS2_TTS_NORM_CACHE_DIR", "/tmp/zonos2-tts-norm")12os.environ.setdefault("GRADIO_SSR_MODE", "false")13os.environ.setdefault("NUMBA_DISABLE_CUDA", "1")14os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")15os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")16 17for cache_dir in (18    os.environ["HF_HOME"],19    os.environ["HF_MODULES_CACHE"],20    os.environ["MPLCONFIGDIR"],21    os.environ["ZONOS2_TTS_NORM_CACHE_DIR"],22):23    Path(cache_dir).mkdir(parents=True, exist_ok=True)24 25 26print("Importing Space runtime dependencies...", flush=True)27import spaces28import gradio as gr29import numpy as np30import torch31 32print("Importing ZONOS2 modules...", flush=True)33from zonos2.message import TTSSamplingParams34from zonos2.tokenizer.textnorm import SERVER_TO_NEMO_LANG, TTSTextNormalizer35from zonos2.tts import TTSLLM36print("Imported ZONOS2 modules.", flush=True)37 38MODEL_ID = "Zyphra/ZONOS2"39SAMPLE_RATE = 4410040LANGUAGES = [41    ("English (US)", "en_us"),42    ("English (UK)", "en_gb"),43    ("French", "fr_fr"),44    ("German", "de"),45    ("Spanish", "es"),46    ("Italian", "it"),47    ("Portuguese (Brazil)", "pt_br"),48    ("Japanese", "ja"),49    ("Mandarin Chinese", "cmn"),50    ("Korean", "ko"),51]52SPEAKING_RATE_BUCKETS = [53    ("Default", "default"),54    ("Very slow", "0"),55    ("Slow", "1"),56    ("Relaxed", "2"),57    ("Natural", "3"),58    ("Bright", "4"),59    ("Fast", "5"),60    ("Very fast", "6"),61    ("Extreme", "7"),62]63 64torch.backends.cuda.matmul.allow_tf32 = True65 66 67def _load_model() -> TTSLLM:68    print(f"Loading {MODEL_ID} for ZeroGPU inference...", flush=True)69    started = time.perf_counter()70    model = TTSLLM(71        model_path=MODEL_ID,72        decode_audio=True,73        cuda_graph_max_bs=0,74        max_running_req=4,75        max_extend_tokens=4096,76        memory_ratio=0.75,77        use_pynccl=False,78    )79    elapsed = time.perf_counter() - started80    print(f"Loaded {MODEL_ID} in {elapsed:.1f}s", flush=True)81    return model82 83 84TTS: TTSLLM | None = None85TEXT_NORMALIZER = TTSTextNormalizer()86TTS_LOCK = threading.Lock()87 88 89def _estimate_duration(*args, **kwargs) -> int:90    max_tokens = kwargs.get("max_tokens")91    if max_tokens is None and len(args) > 4:92        max_tokens = args[4]93    try:94        max_tokens = int(max_tokens)95    except (TypeError, ValueError):96        max_tokens = 76897    base_seconds = 220 if TTS is None else 4598    return min(300, max(60, base_seconds + max_tokens // 12))99 100 101def _pcm_float32_to_wav(audio_bytes: bytes, sample_rate: int = SAMPLE_RATE) -> str:102    audio = np.frombuffer(audio_bytes, dtype=np.float32)103    if audio.size == 0:104        raise gr.Error("The model returned no audio. Try increasing max tokens.")105    audio = np.nan_to_num(audio, nan=0.0, posinf=0.0, neginf=0.0)106    audio = np.clip(audio, -1.0, 1.0)107    audio_i16 = (audio * 32767.0).astype(np.int16)108 109    handle = tempfile.NamedTemporaryFile(suffix=".wav", delete=False)110    handle.close()111    with wave.open(handle.name, "wb") as wav:112        wav.setnchannels(1)113        wav.setsampwidth(2)114        wav.setframerate(sample_rate)115        wav.writeframes(audio_i16.tobytes())116    return handle.name117 118 119def _normalize_text(text: str, language: str, enabled: bool) -> str:120    if not enabled:121        return text122    if language not in SERVER_TO_NEMO_LANG:123        return text124    return TEXT_NORMALIZER.normalize(text, language)125 126 127def _speaking_rate_bucket(value: str) -> int | None:128    if value == "default":129        return None130    return int(value)131 132 133@spaces.GPU(duration=_estimate_duration)134def synthesize(135    text: str,136    language: str,137    text_normalization: bool,138    speaking_rate: str,139    max_tokens: int,140    temperature: float,141    topk: int,142    top_p: float,143    min_p: float,144    repetition_penalty: float,145    seed: int,146):147    text = (text or "").strip()148    if not text:149        raise gr.Error("Enter text to synthesize.")150    if len(text) > 1200:151        raise gr.Error("Keep the prompt under 1200 characters for this Space.")152 153    normalized = _normalize_text(text, language, text_normalization)154    params = TTSSamplingParams(155        temperature=float(temperature),156        topk=int(topk),157        top_p=float(top_p),158        min_p=float(min_p),159        max_tokens=int(max_tokens),160        repetition_window=50,161        repetition_penalty=float(repetition_penalty),162        repetition_codebooks=8,163        seed=None if seed is None or int(seed) < 0 else int(seed),164    )165 166    started = time.perf_counter()167    with TTS_LOCK:168        global TTS169        if TTS is None:170            TTS = _load_model()171        torch.cuda.set_stream(TTS.stream)172        result = TTS.generate_one(173            normalized,174            params,175            decode_audio=True,176            speaking_rate_bucket=_speaking_rate_bucket(speaking_rate),177            quality_buckets=None,178        )179    elapsed = time.perf_counter() - started180 181    wav_path = _pcm_float32_to_wav(result["audio"], result.get("sample_rate", SAMPLE_RATE))182    frames = len(result.get("audio_tokens") or [])183    eos_frame = result.get("eos_frame")184    status = f"Generated {frames} frames in {elapsed:.1f}s"185    if eos_frame is not None:186        status += f" (EOS frame {eos_frame})"187    if normalized != text:188        status += f"\n\nNormalized text: {normalized}"189    return wav_path, status190 191 192CSS = """193main, .gradio-container, .gradio-container > .fillable {194    max-width: 1180px !important;195    margin-inline: auto !important;196}197.compact-status textarea {198    font-family: ui-monospace, SFMono-Regular, Menlo, Consolas, monospace;199}200"""201 202 203with gr.Blocks(title="ZONOS2") as demo:204    gr.Markdown("# ZONOS2")205    with gr.Row():206        with gr.Column(scale=5):207            text = gr.Textbox(208                label="Text",209                value="In the quiet hum of the studio, ZONOS2 turns written words into natural speech.",210                lines=6,211                max_length=1200,212            )213            with gr.Row():214                language = gr.Dropdown(215                    choices=LANGUAGES,216                    value="en_us",217                    label="Language",218                )219                speaking_rate = gr.Dropdown(220                    choices=SPEAKING_RATE_BUCKETS,221                    value="default",222                    label="Speaking rate",223                )224            text_normalization = gr.Checkbox(value=True, label="Text normalization")225            generate = gr.Button("Generate", variant="primary")226        with gr.Column(scale=4):227            audio = gr.Audio(label="Audio", type="filepath", format="wav")228            status = gr.Textbox(229                label="Status",230                lines=5,231                interactive=False,232                elem_classes=["compact-status"],233            )234 235    with gr.Accordion("Sampling", open=False):236        with gr.Row():237            max_tokens = gr.Slider(238                minimum=128,239                maximum=2048,240                step=64,241                value=768,242                label="Max audio tokens",243            )244            seed = gr.Number(value=-1, precision=0, label="Seed (-1 random)")245        with gr.Row():246            temperature = gr.Slider(0.1, 2.0, value=1.15, step=0.05, label="Temperature")247            topk = gr.Slider(1, 512, value=106, step=1, label="Top-k")248        with gr.Row():249            top_p = gr.Slider(0.0, 1.0, value=0.0, step=0.01, label="Top-p")250            min_p = gr.Slider(0.0, 0.5, value=0.18, step=0.01, label="Min-p")251            repetition_penalty = gr.Slider(252                1.0,253                2.0,254                value=1.2,255                step=0.05,256                label="Repetition penalty",257            )258 259    gr.Examples(260        examples=[261            [262                "The first explorers landed just after sunrise, carrying maps, coffee, and impossible optimism.",263                "en_us",264                True,265                "default",266                512,267                1.15,268                106,269                0.0,270                0.18,271                1.2,272                -1,273            ],274            [275                "Le modèle parle avec une voix claire, expressive et naturellement rythmée.",276                "fr_fr",277                True,278                "default",279                512,280                1.15,281                106,282                0.0,283                0.18,284                1.2,285                -1,286            ],287        ],288        inputs=[289            text,290            language,291            text_normalization,292            speaking_rate,293            max_tokens,294            temperature,295            topk,296            top_p,297            min_p,298            repetition_penalty,299            seed,300        ],301    )302 303    generate.click(304        fn=synthesize,305        inputs=[306            text,307            language,308            text_normalization,309            speaking_rate,310            max_tokens,311            temperature,312            topk,313            top_p,314            min_p,315            repetition_penalty,316            seed,317        ],318        outputs=[audio, status],319        api_name="generate",320        concurrency_limit=1,321    )322 323 324if __name__ == "__main__":325    demo.queue(max_size=8, default_concurrency_limit=1).launch(css=CSS)326