import os, re, tempfile import numpy as np import gradio as gr import torch import librosa import soundfile as sf from transformers import AutoTokenizer, AutoModelForCausalLM from neucodec import NeuCodec from phonemizer.backend import EspeakBackend from vinorm import TTSnorm MODEL_ID = "dinhthuan/neutts-air-vi" CODEC_ID = "neuphonic/neucodec" # Optional default reference (put your own files here) DEFAULT_REF_WAV = "assets/reference.wav" DEFAULT_REF_TXT = "assets/reference.txt" SPEECH_START = "<|SPEECH_GENERATION_START|>" SPEECH_END = "<|SPEECH_GENERATION_END|>" device = "cuda" if torch.cuda.is_available() else "cpu" dtype = torch.bfloat16 if device == "cuda" else torch.float32 tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) model = AutoModelForCausalLM.from_pretrained( MODEL_ID, torch_dtype=dtype, trust_remote_code=True, ).to(device) model.eval() codec = NeuCodec.from_pretrained(CODEC_ID).to(device) codec.eval() phonemizer = EspeakBackend(language="vi", preserve_punctuation=True, with_stress=True) def _phones_vi(text: str) -> str: # same normalization style as model card example t = TTSnorm(text, punc=False, unknown=True, lower=False, rule=False) return phonemizer.phonemize([t])[0] def _encode_ref_16k(ref_wav_path: str) -> torch.Tensor: # model card encodes reference audio at 16kHz mono wav, _ = librosa.load(ref_wav_path, sr=16000, mono=True) wav = torch.from_numpy(wav).float().unsqueeze(0).unsqueeze(0).to(device) # (1,1,T) with torch.no_grad(): codes = codec.encode_code(audio_or_path=wav).squeeze(0).squeeze(0).detach().cpu() return codes def _extract_codes(text: str) -> list[int]: # take codes between generation tags if present if SPEECH_START in text and SPEECH_END in text: text = text.split(SPEECH_START, 1)[1].split(SPEECH_END, 1)[0] return [int(x) for x in re.findall(r"<\|speech_(\d+)\|>", text)] @torch.inference_mode() def tts(text: str, ref_audio_path: str | None, ref_text: str | None, max_new_tokens: int): text = (text or "").strip() if not text: raise gr.Error("Hãy nhập văn bản tiếng Việt.") # choose reference if ref_audio_path: wav_path = ref_audio_path rt = (ref_text or "").strip() if not rt: raise gr.Error("Bạn đã upload reference audio thì cần nhập reference text (đúng nội dung audio).") else: # fallback to default assets if not (os.path.exists(DEFAULT_REF_WAV) and os.path.exists(DEFAULT_REF_TXT)): raise gr.Error( "Chưa có reference. Hãy upload ref audio + ref text, " "hoặc thêm assets/reference.wav và assets/reference.txt vào repo." ) wav_path = DEFAULT_REF_WAV rt = open(DEFAULT_REF_TXT, "r", encoding="utf-8").read().strip() # phonemize phones = _phones_vi(text) ref_phones = _phones_vi(rt) # encode reference audio to speech codes ref_codes = _encode_ref_16k(wav_path) codes_str = "".join([f"<|speech_{i}|>" for i in ref_codes.tolist()]) combined_phones = ref_phones + " " + phones # prompt format follows model card chat = ( "user: Convert the text to speech:" f"<|TEXT_PROMPT_START|>{combined_phones}<|TEXT_PROMPT_END|>\n" f"assistant:{SPEECH_START}{codes_str}" ) input_ids = tokenizer.encode(chat, return_tensors="pt").to(device) speech_end_id = tokenizer.convert_tokens_to_ids(SPEECH_END) out = model.generate( input_ids, max_new_tokens=int(max_new_tokens), temperature=1.0, top_k=50, eos_token_id=speech_end_id, pad_token_id=tokenizer.eos_token_id, ) out_text = tokenizer.decode(out[0], skip_special_tokens=False) all_codes = _extract_codes(out_text) # remove the prefix ref_codes if present gen_codes = all_codes[len(ref_codes):] if len(all_codes) > len(ref_codes) else all_codes if len(gen_codes) < 10: raise gr.Error("Không trích xuất được speech codes. Hãy thử ref audio rõ hơn hoặc text ngắn hơn.") codes_tensor = torch.tensor(gen_codes, dtype=torch.long).view(1, 1, -1).to(device) audio = codec.decode_code(codes_tensor).detach().cpu().numpy()[0, 0, :] audio = np.clip(audio, -1.0, 1.0) tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) # model card indicates output sample rate 24kHz sf.write(tmp.name, audio, 24000) return tmp.name with gr.Blocks(title="Vietnamese TTS (NeuTTS-Air finetune)") as demo: gr.Markdown("## Vietnamese TTS – dinhthuan/neutts-air-vi\nNhập tiếng Việt → Xuất âm thanh (WAV 24kHz).") text_in = gr.Textbox(label="Văn bản tiếng Việt", lines=4, value="Xin chào, đây là mô hình TTS tiếng Việt.") with gr.Row(): ref_audio = gr.Audio(label="Reference audio (3–10s, WAV)", type="filepath") ref_text = gr.Textbox(label="Reference text (đúng nội dung của ref audio)", lines=2) max_tok = gr.Slider(256, 3072, value=1536, step=128, label="max_new_tokens") btn = gr.Button("Tạo giọng nói") out_audio = gr.Audio(label="Kết quả", type="filepath") btn.click(tts, inputs=[text_in, ref_audio, ref_text, max_tok], outputs=out_audio) demo.queue().launch()