CoolFace
Apppublic

flaviusburca/DramaboxTTS

sourceHugging Faceotherupdated 4mo agoView on Hugging Face
0likes
app.py411 linesDownload Raw Back to root
1#!/usr/bin/env python32"""DramaBox TTS — Voice Gallery (ZeroGPU)"""3 4import json5import logging6import os7import sys8import tempfile9 10import gradio as gr11import requests12import soundfile as sf13import spaces14 15_DIR = os.path.dirname(os.path.abspath(__file__))16sys.path.insert(0, os.path.join(_DIR, "src"))17from inference_server import TTSServer  # noqa: E40218from model_downloader import get_all_paths  # noqa: E40219from duration_estimator import estimate_speech_duration  # noqa: E40220from text_chunker import chunk_prompt_for_duration  # noqa: E40221import higgs_backend  # noqa: E40222import asr_backend  # noqa: E40223 24logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")25 26# ── Voice backends ─────────────────────────────────────────────────────────────27# Both backends load eagerly at startup and stay warm — ZeroGPU's H200 (141 GB)28# comfortably holds DramaBox (22B 4-bit transformer + 12B 4-bit Gemma) and29# Higgs Audio v3 (4B bf16) side by side, so switching backends mid-session30# never pays a model-load tax.31BACKEND_DRAMABOX = "DramaBox (LTX-2)"32BACKEND_HIGGS = "Higgs Audio v3 (4B)"33BACKENDS = [BACKEND_DRAMABOX, BACKEND_HIGGS]34 35# ── Voices ─────────────────────────────────────────────────────────────────────36with open(os.path.join(_DIR, "voices.json"), encoding="utf-8") as _f:37    VOICES = json.load(_f)38 39LANGUAGES = ["All"] + sorted({v.get("language", "") for v in VOICES if v.get("language")})40GENDERS   = ["All", "female", "male"]41PER_PAGE  = 2042 43logging.info(f"Loaded {len(VOICES):,} voices")44 45 46def _filter(search, lang, gender, accent):47    s = (search or "").lower()48    return [49        v for v in VOICES50        if (lang   == "All" or v.get("language") == lang)51        and (gender == "All" or v.get("gender")   == gender)52        and (accent == "All" or v.get("accent")   == accent)53        and (not s  or s in v.get("name", "").lower()54                    or s in (v.get("description") or "").lower())55    ]56 57 58def _accents_for(lang):59    pool = VOICES if lang == "All" else [v for v in VOICES if v.get("language") == lang]60    return ["All"] + sorted({v.get("accent", "") for v in pool if v.get("accent")})61 62 63# ── Model ──────────────────────────────────────────────────────────────────────64logging.info("Downloading model weights…")65PATHS = get_all_paths()66tts = TTSServer(67    checkpoint=PATHS["transformer"],68    full_checkpoint=PATHS["audio_components"],69    gemma_root=PATHS["gemma_root"],70    device="cuda",71    compile_model=False,72    bnb_4bit=True,73)74logging.info("TTSServer ready.")75 76higgs_backend.load()77asr_backend.load()78 79 80@spaces.GPU(duration=15)81def on_generate(backend, prompt, preview_url, cfg, stg, dur_mult, seed,82                max_c, tgt_c, ref_text, temperature, top_p, top_k, max_new_tok, higgs_seed,83                progress=gr.Progress()):84    if not (prompt or "").strip():85        raise gr.Error("Prompt is empty.")86 87    ref_path = None88    if preview_url:89        r = requests.get(preview_url, timeout=30)90        r.raise_for_status()91        tmp = tempfile.NamedTemporaryFile(suffix=".mp3", delete=False)92        tmp.write(r.content)93        tmp.close()94        ref_path = tmp.name95 96    try:97        if backend == BACKEND_HIGGS:98            progress(0.5, desc="Generating with Higgs Audio v3…")99            waveform, sr = higgs_backend.generate(100                prompt.strip(), voice_ref=ref_path, reference_text=ref_text,101                temperature=float(temperature), top_p=float(top_p),102                top_k=int(top_k), max_new_tokens=int(max_new_tok), seed=int(higgs_seed),103            )104            out = tempfile.mktemp(suffix=".wav", prefix="higgs_", dir="/tmp")105            sf.write(out, waveform, sr)106            return out107 108        out = tempfile.mktemp(suffix=".wav", prefix="dramabox_", dir="/tmp")109 110        def _prog(idx, total, est_s):111            progress(idx / total if total > 1 else 0.,112                     desc=f"Chunk {idx+1}/{total}" if total > 1 else f"Generating ~{est_s:.1f}s…")113 114        tts.generate_to_file(115            prompt=prompt.strip(), output=out, voice_ref=ref_path,116            cfg_scale=float(cfg), stg_scale=float(stg),117            duration_multiplier=float(dur_mult), seed=int(seed),118            max_chunk_duration=float(max_c), target_chunk_duration=float(tgt_c),119            progress_callback=_prog,120        )121        return out122    finally:123        if ref_path and os.path.exists(ref_path):124            os.unlink(ref_path)125 126 127# ── CSS ─────────────────────────────────────────────────────────────────────────128CSS = """129/* card grid */130.card-grid { display: grid; grid-template-columns: repeat(4, 1fr); gap: 12px; }131@media (max-width: 1200px) { .card-grid { grid-template-columns: repeat(3, 1fr); } }132@media (max-width: 800px)  { .card-grid { grid-template-columns: repeat(2, 1fr); } }133 134/* individual card — scoped inside the Gradio column */135.voice-card { background: #16161e !important; border: 1px solid #2a2a3a !important;136              border-radius: 10px !important; padding: 14px !important; height: 100% !important; }137.voice-card:hover { border-color: #ff6b35 !important; }138 139/* card header line */140.card-header { display: flex; align-items: flex-start; gap: 8px; margin-bottom: 6px; }141.badge-f { background: #3d0e3d; color: #e080e0; font-size: 11px; font-weight: 700;142           padding: 2px 7px; border-radius: 4px; white-space: nowrap; }143.badge-m { background: #0e1e3d; color: #80a8e0; font-size: 11px; font-weight: 700;144           padding: 2px 7px; border-radius: 4px; white-space: nowrap; }145.card-name { font-size: 13px; font-weight: 600; color: #dde0f0; line-height: 1.35; }146 147/* tags row */148.card-tags { display: flex; flex-wrap: wrap; gap: 4px; margin-bottom: 4px; }149.card-tags span { font-size: 10px; padding: 2px 6px; border-radius: 3px; }150.t-lang { background: #1e3a1e; color: #88cc88; }151.t-acc  { background: #1e2a3a; color: #88a8cc; }152.t-age  { background: #2a1e2a; color: #aa88aa; }153 154/* description */155.card-desc { font-size: 11px; color: #5a5a80; line-height: 1.4; margin-bottom: 4px; }156 157/* "Use this voice" button override */158.use-btn { background: #ff6b35 !important; border: none !important;159           font-weight: 700 !important; }160.use-btn:hover { background: #ff8755 !important; }161 162/* selected voice banner */163.sel-banner { background: #0d1a0d; border: 1px solid #2a4a2a; border-radius: 8px;164              padding: 10px 14px; margin: 6px 0; }165 166/* pagination */167.pager-row { display: flex; align-items: center; gap: 12px; padding: 8px 0; }168 169"""170 171# ── UI ──────────────────────────────────────────────────────────────────────────172with gr.Blocks(title="DramaBox TTS", css=CSS, analytics_enabled=False) as app:173 174    gr.Markdown(f"# 🎭 DramaBox TTS\nBrowse **{len(VOICES):,} voices**. Hit ▶ to preview, then **Use this voice** to generate.")175 176    # Filters177    with gr.Row():178        search_in  = gr.Textbox(placeholder="Search by name or description…", label="Search", scale=3)179        lang_in    = gr.Dropdown(LANGUAGES, value="All", label="Language", scale=2)180        gender_in  = gr.Radio(GENDERS, value="All", label="Gender", scale=2)181        accent_in  = gr.Dropdown(["All"], value="All", label="Accent", scale=2)182 183    result_md = gr.Markdown("")184 185    # ── Fixed card grid (PER_PAGE slots) ───────────────────────────────────────186    # Build PER_PAGE card slots; each slot has HTML header, Audio, Use button.187    # Slots are hidden when a page has fewer voices than PER_PAGE.188    card_rows = []   # gr.Column slots (show/hide)189    card_html = []   # gr.HTML — full card content incl. <audio> tag190    card_btns = []   # gr.Button — "Use this voice"191 192    page_voices = gr.State([])   # voice dicts on the current page193 194    COLS = 4195    for r_idx in range((PER_PAGE + COLS - 1) // COLS):196        with gr.Row():197            for c_idx in range(COLS):198                slot = r_idx * COLS + c_idx199                if slot >= PER_PAGE:200                    break201                with gr.Column(elem_classes=["voice-card"]) as col:202                    html = gr.HTML("")203                    btn  = gr.Button("✅ Use this voice", size="sm",204                                     elem_classes=["use-btn"])205                card_html.append(html)206                card_btns.append(btn)207                card_rows.append(col)208 209    # Pagination210    with gr.Row(elem_classes=["pager-row"]):211        prev_btn  = gr.Button("← Prev", size="sm", interactive=False)212        page_info = gr.Markdown("", elem_classes=["pager-info"])213        next_btn  = gr.Button("Next →", size="sm", interactive=False)214 215    # Selected voice banner216    with gr.Row(visible=False, elem_classes=["sel-banner"]) as sel_row:217        with gr.Column(scale=2):218            sel_md    = gr.Markdown("**No voice selected**")219        with gr.Column(scale=3):220            sel_audio = gr.Audio(label="Selected voice preview", type="filepath",221                                 interactive=False, show_download_button=False)222    sel_url = gr.State(None)223 224    # Generation225    gr.Markdown("---\n## ✏️ Write your scene prompt")226    gr.Markdown(227        "Dialogue inside `\"double quotes\"`, stage directions outside. "228        "Phonetics (`\"Hahaha\"`, `\"Mmm\"`) inside quotes; actions (`She sighs.`) outside.\n\n"229        "*(The prompt format above is for the DramaBox backend — Higgs Audio v3 just reads "230        "plain text aloud in the cloned/selected voice.)*"231    )232    backend_in = gr.Radio(233        BACKENDS, value=BACKEND_DRAMABOX, label="Voice backend",234        info="DramaBox: expressive scene-prompt TTS with stage directions (LTX-2). "235             "Higgs Audio v3: fast zero-shot voice cloning (4B).",236    )237    with gr.Row():238        with gr.Column(scale=3):239            prompt_box = gr.Textbox(240                label="Scene prompt", lines=5,241                placeholder='A woman says warmly, "Hey, just saying hi — hope you\'re doing well!"\n'242                            'She laughs softly, "Hehehe, we should hang out soon."',243            )244            gen_btn = gr.Button("🎙️ Generate", variant="primary", size="lg")245        with gr.Column(scale=2):246            with gr.Group(visible=True) as dramabox_settings:247                with gr.Accordion("Settings", open=False):248                    cfg_s  = gr.Slider(1., 10., 2.5, step=.5,  label="CFG scale")249                    stg_s  = gr.Slider(0., 5.,  1.5, step=.5,  label="STG scale")250                    dur_s  = gr.Slider(.8, 2.,  1.1, step=.05, label="Duration ×")251                    seed_n = gr.Number(42, precision=0, label="Seed")252                with gr.Accordion("Chunking", open=False):253                    max_c  = gr.Slider(20., 60., 45., step=1., label="Max chunk (s)")254                    tgt_c  = gr.Slider(15., 50., 37., step=1., label="Target chunk (s)")255            with gr.Group(visible=False) as higgs_settings:256                with gr.Accordion("Settings", open=False):257                    ref_text_in = gr.Textbox(258                        label="Reference transcript (auto-filled on selection, improves cloning)",259                        lines=2, placeholder="Auto-transcribed from the selected voice's preview clip — edit or clear as needed.",260                    )261                    temperature_s = gr.Slider(0., 1.5, .7,   step=.05, label="Temperature")262                    top_p_s       = gr.Slider(.1, 1.,  .95,  step=.01, label="Top-p")263                    top_k_s       = gr.Slider(0, 1026, 50,   step=1,   label="Top-k (0 = off)")264                    max_tok_s     = gr.Slider(64, 4096, 2048, step=64, label="Max new tokens")265                    higgs_seed_n  = gr.Number(-1, precision=0, label="Seed (-1 = random)")266    audio_out = gr.Audio(label="Generated audio", type="filepath")267 268    def on_backend_change(backend):269        is_higgs = backend == BACKEND_HIGGS270        return gr.update(visible=not is_higgs), gr.update(visible=is_higgs)271 272    backend_in.change(on_backend_change, backend_in, [dramabox_settings, higgs_settings])273 274    # ── Page state ─────────────────────────────────────────────────────────────275    page_state = gr.State(1)276 277    # ── Helper: build all card + pagination outputs from a voice list + page ────278    def _all_updates(filtered, page):279        total       = len(filtered)280        total_pages = max(1, (total + PER_PAGE - 1) // PER_PAGE)281        page        = max(1, min(page, total_pages))282        chunk       = filtered[(page - 1) * PER_PAGE : page * PER_PAGE]283 284        html_updates, vis_updates = [], []285        for i in range(PER_PAGE):286            if i < len(chunk):287                v     = chunk[i]288                g     = v.get("gender", "")289                badge = f'<span class="badge-{"f" if g=="female" else "m"}">{"♀" if g=="female" else "♂"}</span>'290                name  = v.get("name", "Unknown")291                lt, at, ag = v.get("language","?"), v.get("accent","?"), v.get("age","?")292                desc  = (v.get("description") or "")[:100]293                src   = v.get("preview_url", "")294                html  = (295                    f'<div class="card-header">{badge}'296                    f'<span class="card-name">{name}</span></div>'297                    f'<div class="card-tags">'298                    f'<span class="t-lang">{lt}</span>'299                    f'<span class="t-acc">{at}</span>'300                    f'<span class="t-age">{ag}</span></div>'301                    + (f'<p class="card-desc">{desc}</p>' if desc else "")302                    + f'<audio controls preload="none" src="{src}" style="width:100%;height:32px;margin-top:4px"></audio>'303                )304                html_updates.append(gr.update(value=html))305                vis_updates.append(gr.update(visible=True))306            else:307                html_updates.append(gr.update(value=""))308                vis_updates.append(gr.update(visible=False))309 310        return (311            html_updates + vis_updates +312            [gr.update(value=f"**{total:,}** voices found"),313             gr.update(value=f"Page **{page}** / {total_pages}"),314             gr.update(interactive=page > 1),315             gr.update(interactive=page < total_pages),316             chunk, page]317        )318 319    _gallery_outputs = (320        card_html + card_rows +321        [result_md, page_info, prev_btn, next_btn, page_voices, page_state]322    )323 324    # ── Filter change → reset to page 1 ────────────────────────────────────────325    def on_filter(s, l, g, a):326        filtered = _filter(s, l, g, a)327        return _all_updates(filtered, 1)328 329    def on_lang(l):330        return gr.Dropdown(choices=_accents_for(l), value="All")331 332    lang_in.change(on_lang, lang_in, accent_in)333 334    for inp in [search_in, lang_in, gender_in, accent_in]:335        inp.change(on_filter,336                   [search_in, lang_in, gender_in, accent_in],337                   _gallery_outputs)338 339    # ── Pagination ──────────────────────────────────────────────────────────────340    def on_prev(s, l, g, a, pg):341        return _all_updates(_filter(s, l, g, a), int(pg) - 1)342 343    def on_next(s, l, g, a, pg):344        return _all_updates(_filter(s, l, g, a), int(pg) + 1)345 346    prev_btn.click(on_prev, [search_in, lang_in, gender_in, accent_in, page_state], _gallery_outputs)347    next_btn.click(on_next, [search_in, lang_in, gender_in, accent_in, page_state], _gallery_outputs)348 349    # ── "Use this voice" buttons ────────────────────────────────────────────────350    def _make_use_handler(slot_idx):351        def handler(voices):352            if slot_idx >= len(voices):353                return gr.update(), gr.update(), gr.update(), gr.update(visible=False), None354            v       = voices[slot_idx]355            name    = v.get("name", "Unknown")356            preview = v.get("preview_url", "")357            tmp     = None358            if preview:359                try:360                    r = requests.get(preview, timeout=15)361                    r.raise_for_status()362                    f = tempfile.NamedTemporaryFile(suffix=".mp3", delete=False)363                    f.write(r.content)364                    f.close()365                    tmp = f.name366                except Exception as e:367                    logging.warning(f"Preview download failed: {e}")368            return (369                gr.update(value=f"**Selected:** {name}"),370                gr.update(value=tmp),371                gr.update(visible=True),372                preview,373            )374        return handler375 376    for i, btn in enumerate(card_btns):377        btn.click(378            _make_use_handler(i),379            inputs=[page_voices],380            outputs=[sel_md, sel_audio, sel_row, sel_url],381        )382 383    # Auto-transcribe the selected voice's preview clip on CPU (Whisper) so384    # "Reference transcript" is pre-filled for Higgs Audio v3 cloning — the385    # user can still edit or clear it before generating.386    sel_audio.change(asr_backend.transcribe, inputs=[sel_audio], outputs=[ref_text_in])387 388    # ── Generate ────────────────────────────────────────────────────────────────389    gen_btn.click(390        on_generate,391        [backend_in, prompt_box, sel_url, cfg_s, stg_s, dur_s, seed_n, max_c, tgt_c,392         ref_text_in, temperature_s, top_p_s, top_k_s, max_tok_s, higgs_seed_n],393        [audio_out],394    )395 396    # ── Initial load ────────────────────────────────────────────────────────────397    app.load(398        lambda: _all_updates(VOICES, 1),399        outputs=_gallery_outputs,400    )401 402 403if __name__ == "__main__":404    port = int(os.environ.get("GRADIO_SERVER_PORT", "7860"))405    app.queue(max_size=10).launch(406        server_name="0.0.0.0", server_port=port,407        share=os.environ.get("GRADIO_SHARE", "1") == "1",408        ssr_mode=False,409        show_api=False,410    )411