Spaces:
Sleeping
Sleeping
| 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)] | |
| 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() | |