import os import random import gc import gradio as gr import spaces import torch import torchaudio from transformers import ( AutoModelForCausalLM, AutoTokenizer, AutoProcessor, MoonshineForConditionalGeneration, ) os.environ["TOKENIZERS_PARALLELISM"] = "false" MODEL_REPO = "multimodalart/higgs-audio-v3-tts-4b-transformers" ASR_REPO = "UsefulSensors/moonshine-base" tokenizer = AutoTokenizer.from_pretrained( MODEL_REPO, trust_remote_code=True, ) model = AutoModelForCausalLM.from_pretrained( MODEL_REPO, trust_remote_code=True, dtype=torch.bfloat16, ).to("cuda").eval() SAMPLE_RATE = int(getattr(model.config, "sample_rate", 24000)) asr_processor = AutoProcessor.from_pretrained(ASR_REPO) asr_model = MoonshineForConditionalGeneration.from_pretrained(ASR_REPO).eval() def load_audio_for_higgs(path: str): """ Official model-card pattern uses torchaudio.load: ref, sr = torchaudio.load("reference.wav") reference_audio=ref reference_sample_rate=sr torchaudio.load returns [channels, time], which is what we pass through. """ wav, sr = torchaudio.load(path) return wav.float(), int(sr) def transcribe(reference_audio): """ Runs on CPU. This is only for auto-filling the reference transcript box. """ if reference_audio is None: return "" try: wav, sr = torchaudio.load(reference_audio) wav = wav.mean(dim=0, keepdim=True) if sr != 16000: wav = torchaudio.functional.resample(wav, sr, 16000) audio_np = wav.squeeze(0).numpy() inputs = asr_processor( audio_np, sampling_rate=16000, return_tensors="pt", ) with torch.inference_mode(): generated_ids = asr_model.generate(**inputs) text = asr_processor.decode( generated_ids[0], skip_special_tokens=True, ) return text.strip() except Exception as e: print(f"[transcribe] {type(e).__name__}: {e}") return "" def estimate_duration(text, reference_audio, reference_text, temperature, top_p, top_k, max_new_tokens, seed): try: tokens = int(max_new_tokens) except Exception: tokens = 1024 if tokens <= 512: return 60 if tokens <= 1024: return 90 if tokens <= 2048: return 120 return 180 @spaces.GPU(duration=estimate_duration, size="xlarge") def synthesize(text, reference_audio, reference_text, temperature, top_p, top_k, max_new_tokens, seed): text = (text or "").strip() if not text: raise gr.Error("Please enter text to synthesize.") seed = int(seed) if seed >= 0: random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) kwargs = { "temperature": float(temperature), "max_new_tokens": int(max_new_tokens), } top_p = float(top_p) top_k = int(top_k) if top_p < 1.0: kwargs["top_p"] = top_p if top_k > 0: kwargs["top_k"] = top_k if reference_audio is not None: ref_wav, ref_sr = load_audio_for_higgs(reference_audio) kwargs["reference_audio"] = ref_wav kwargs["reference_sample_rate"] = ref_sr reference_text = (reference_text or "").strip() if reference_text: kwargs["reference_text"] = reference_text try: with torch.inference_mode(): wav = model.generate_speech( text, tokenizer, **kwargs, ) if wav is None or wav.numel() == 0: raise gr.Error("Generation returned no audio. Try shorter text or fewer max new tokens.") wav = wav.detach().cpu().float().numpy() gc.collect() torch.cuda.empty_cache() return (SAMPLE_RATE, wav) except gr.Error: raise except Exception as e: print(f"[synthesize] {type(e).__name__}: {e}") raise gr.Error(f"Generation failed: {type(e).__name__}: {e}") with gr.Blocks( title="Higgs Audio v3 TTS", ) as demo: gr.Markdown( "# Higgs Audio v3 TTS\n" "Zero-shot text-to-speech and voice cloning with " "[`multimodalart/higgs-audio-v3-tts-4b-transformers`](https://huggingface.co/multimodalart/higgs-audio-v3-tts-4b-transformers)." ) with gr.Row(): with gr.Column(): text = gr.Textbox( label="Text to synthesize", placeholder="Type what you want the voice to say…", lines=4, ) reference_audio = gr.Audio( label="Reference voice (optional — for cloning)", type="filepath", ) reference_text = gr.Textbox( label="Reference transcript (auto-filled on upload, improves cloning)", lines=2, ) with gr.Accordion("Advanced settings", open=False): temperature = gr.Slider( minimum=0.0, maximum=1.5, value=0.7, step=0.05, label="Temperature", ) top_p = gr.Slider( minimum=0.1, maximum=1.0, value=0.95, step=0.01, label="Top-p", ) top_k = gr.Slider( minimum=0, maximum=1026, value=50, step=1, label="Top-k (0 = off)", ) max_new_tokens = gr.Slider( minimum=64, maximum=4096, value=1024, step=64, label="Max new tokens", ) seed = gr.Number( value=-1, label="Seed (-1 = random)", precision=0, ) run_btn = gr.Button("Generate speech", variant="primary") with gr.Column(): output_audio = gr.Audio( label="Generated speech", type="numpy", ) gr.Examples( examples=[ [ "Higgs Audio version three, now running on plain transformers.", None, "", ], [ "The quick brown fox jumps over the lazy dog.", None, "", ], ], inputs=[text, reference_audio, reference_text], ) reference_audio.change( fn=transcribe, inputs=reference_audio, outputs=reference_text, api_name="transcribe", ) run_btn.click( fn=synthesize, inputs=[ text, reference_audio, reference_text, temperature, top_p, top_k, max_new_tokens, seed, ], outputs=output_audio, api_name="synthesize", ) demo.queue().launch( theme=gr.themes.Citrus(), show_error=True, )