flaviusburca/DramaboxTTS
0
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 