Download generate_tts.py from FraPiz/xtts-v2-moldovan-romanian: direct link, hf CLI and curl.
- Browser
- Download file 11.2 kB
-
https://huggingface.co/FraPiz/xtts-v2-moldovan-romanian/resolve/main/generate_tts.py
- Command line
-
hf download hf://FraPiz/xtts-v2-moldovan-romanian/generate_tts.py
-
curl -L -o generate_tts.py https://huggingface.co/FraPiz/xtts-v2-moldovan-romanian/resolve/main/generate_tts.py
11.2 kB
| """ | |
| XTTS-v2 Moldovan TTS — Generare audio din text | |
| Folosește modelul fine-tunat pe grai moldovenesc. | |
| Prima rulare: extrage și salvează latentele vocii de referință (~10s). | |
| Rulările următoare: încarcă latentele salvate (~instant). | |
| """ | |
| import os | |
| import sys | |
| import time | |
| import torch | |
| import soundfile as sf | |
| import numpy as np | |
| # ══════════════════════════════════════════════════════════════ | |
| # TEXT DE GENERAT (modifică aici) | |
| # ══════════════════════════════════════════════════════════════ | |
| TEXT = """ | |
| Bună ziua! Eu sunt un model de sinteză vocală antrenat pe grai moldovenesc. | |
| """ | |
| # ══════════════════════════════════════════════════════════════ | |
| # CONFIGURARE | |
| # ══════════════════════════════════════════════════════════════ | |
| SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__)) | |
| # Căi model | |
| MODEL_DIR = SCRIPT_DIR | |
| CHECKPOINT_PATH = os.path.join(MODEL_DIR, "best_model.pth") | |
| CONFIG_PATH = os.path.join(MODEL_DIR, "config.json") | |
| VOCAB_PATH = os.path.join(MODEL_DIR, "vocab.json") | |
| # Voce de referință | |
| REFERENCE_AUDIO = os.path.join(SCRIPT_DIR, "reference_voice", "reference.wav") | |
| CACHED_LATENTS = os.path.join(SCRIPT_DIR, "reference_voice", "cached_latents.pth") | |
| # Output | |
| OUTPUT_DIR = os.path.join(SCRIPT_DIR, "output") | |
| os.makedirs(OUTPUT_DIR, exist_ok=True) | |
| # Parametri generare | |
| TEMPERATURE = 0.7 | |
| TOP_P = 0.7 | |
| TOP_K = 30 | |
| LENGTH_PENALTY = 0.8 | |
| REPETITION_PENALTY = 10.0 | |
| GPT_COND_LEN = 6 # Secunde de audio referință pentru condiționare | |
| # Cedilla → comma-below | |
| CEDILLA_TO_COMMA = str.maketrans({ | |
| "\u015f": "\u0219", # ş -> ș | |
| "\u0163": "\u021b", # ţ -> ț | |
| "\u015e": "\u0218", # Ş -> Ș | |
| "\u0162": "\u021a", # Ţ -> Ț | |
| }) | |
| # ══════════════════════════════════════════════════════════════ | |
| # PATCH TOKENIZER ROMÂNESC | |
| # ══════════════════════════════════════════════════════════════ | |
| def patch_tokenizer_for_romanian(): | |
| """Aplică patch pentru suport limba română în tokenizer-ul XTTS.""" | |
| import re | |
| import TTS | |
| tts_dir = os.path.dirname(TTS.__file__) | |
| tokenizer_file = os.path.join(tts_dir, "tts", "layers", "xtts", "tokenizer.py") | |
| if not os.path.exists(tokenizer_file): | |
| print(f"WARN: tokenizer.py negăsit la {tokenizer_file}") | |
| return | |
| with open(tokenizer_file, "r", encoding="utf-8") as f: | |
| content = f.read() | |
| original = content | |
| # Adaugă 'ro' la listele de limbi | |
| if '"ro"' not in content and "'ro'" not in content: | |
| content = re.sub( | |
| r"""(['"])zh-cn\1(\s*)\]""", | |
| lambda m: f'{m.group(1)}zh-cn{m.group(1)}, {m.group(1)}ro{m.group(1)}{m.group(2)}]', | |
| content | |
| ) | |
| # Safe .get() replacements | |
| for old, new in { | |
| "self.abbreviations[lang]": "self.abbreviations.get(lang, [])", | |
| "self.symbols[lang]": "self.symbols.get(lang, [])", | |
| "_abbreviations[lang]": "_abbreviations.get(lang, [])", | |
| "_symbols_multilingual[lang]": "_symbols_multilingual.get(lang, [])", | |
| 'and_equivalents[lang]': 'and_equivalents.get(lang, "and")', | |
| '_ordinal_re[lang]': '_ordinal_re.get(lang, re.compile(r"(?!x)x"))', | |
| }.items(): | |
| if old in content: | |
| content = content.replace(old, new) | |
| # Cedilla normalization block | |
| if "_CEDILLA_TO_COMMA" not in content: | |
| import_match = re.search(r'^import re\s*$', content, re.MULTILINE) | |
| if import_match: | |
| cedilla_block = ( | |
| '\n# Cedilla -> comma-below normalization for Romanian\n' | |
| '_CEDILLA_TO_COMMA = str.maketrans({\n' | |
| ' "\u015f": "\u0219",\n' | |
| ' "\u0163": "\u021b",\n' | |
| ' "\u015e": "\u0218",\n' | |
| ' "\u0162": "\u021a",\n' | |
| '})\n' | |
| ) | |
| pos = import_match.end() | |
| content = content[:pos] + cedilla_block + content[pos:] | |
| # Inject cedilla normalization in multilingual_cleaners | |
| if "_CEDILLA_TO_COMMA" in content: | |
| mc_match = re.search(r'(def multilingual_cleaners\([^)]*\):\s*\n)', content) | |
| if mc_match: | |
| func_start = mc_match.end() | |
| if "_CEDILLA_TO_COMMA" not in content[func_start:func_start + 500]: | |
| injection = ' if lang == "ro":\n text = text.translate(_CEDILLA_TO_COMMA)\n' | |
| content = content[:func_start] + injection + content[func_start:] | |
| # Add 'ro' to _ordinal_re dict | |
| if "_ordinal_re" in content: | |
| ord_match = re.search(r'_ordinal_re\s*=\s*\{', content) | |
| if ord_match: | |
| brace_start = content.index('{', ord_match.start()) | |
| depth = 0 | |
| brace_end = brace_start | |
| for ci in range(brace_start, len(content)): | |
| if content[ci] == '{': | |
| depth += 1 | |
| elif content[ci] == '}': | |
| depth -= 1 | |
| if depth == 0: | |
| brace_end = ci | |
| break | |
| ordinal_dict = content[brace_start:brace_end + 1] | |
| if '"ro"' not in ordinal_dict and "'ro'" not in ordinal_dict: | |
| new_dict = ordinal_dict[:-1] + ' "ro": re.compile(r"([0-9]+)\\.(?=\\s|$)"),\n}' | |
| content = content[:brace_start] + new_dict + content[brace_end + 1:] | |
| # Add 'ro' to preprocess_text | |
| for old_pat, new_pat in [ | |
| ('"pt", "ru"', '"pt", "ro", "ru"'), | |
| ("'pt', 'ru'", "'pt', 'ro', 'ru'"), | |
| ]: | |
| if old_pat in content: | |
| content = content.replace(old_pat, new_pat, 1) | |
| break | |
| if content != original: | |
| with open(tokenizer_file, "w", encoding="utf-8") as f: | |
| f.write(content) | |
| # Reload affected modules | |
| mods_to_reload = [k for k in sys.modules if 'xtts' in k.lower() or 'tokenizer' in k.lower()] | |
| for mod in mods_to_reload: | |
| del sys.modules[mod] | |
| print(" Tokenizer patch românesc aplicat!") | |
| else: | |
| print(" Tokenizer deja patch-uit.") | |
| # ══════════════════════════════════════════════════════════════ | |
| # MAIN | |
| # ══════════════════════════════════════════════════════════════ | |
| def main(): | |
| t_start = time.time() | |
| # Verificări | |
| if not os.path.exists(CHECKPOINT_PATH): | |
| print(f"ERROR: best_model.pth nu există la {CHECKPOINT_PATH}") | |
| sys.exit(1) | |
| if not os.path.exists(REFERENCE_AUDIO): | |
| print(f"ERROR: Audio de referință nu există la {REFERENCE_AUDIO}") | |
| print(f"Copiază fișierul WAV în: {os.path.dirname(REFERENCE_AUDIO)}") | |
| sys.exit(1) | |
| # GPU check | |
| if torch.cuda.is_available(): | |
| gpu_name = torch.cuda.get_device_name(0) | |
| vram = torch.cuda.get_device_properties(0).total_memory / 1024**3 | |
| print(f"GPU: {gpu_name} ({vram:.1f} GB VRAM)") | |
| else: | |
| print("WARN: GPU nu a fost detectat! Generarea va fi foarte lentă.") | |
| # Patch tokenizer | |
| print("\nAplicare patch tokenizer românesc...") | |
| patch_tokenizer_for_romanian() | |
| # Patch torchaudio.load pentru a folosi soundfile (evită FFmpeg/torchcodec pe Windows) | |
| import torchaudio | |
| def _soundfile_load(filepath, **kwargs): | |
| data, sr = sf.read(filepath, dtype='float32') | |
| if data.ndim == 1: | |
| data = data[np.newaxis, :] | |
| else: | |
| data = data.T | |
| return torch.from_numpy(data), sr | |
| torchaudio.load = _soundfile_load | |
| # Încărcare model | |
| print("\nÎncărcare model XTTS-v2...") | |
| t_load = time.time() | |
| from TTS.tts.configs.xtts_config import XttsConfig | |
| from TTS.tts.models.xtts import Xtts | |
| config = XttsConfig() | |
| config.load_json(CONFIG_PATH) | |
| model = Xtts.init_from_config(config) | |
| model.load_checkpoint( | |
| config, | |
| checkpoint_path=CHECKPOINT_PATH, | |
| vocab_path=VOCAB_PATH, | |
| use_deepspeed=False, | |
| ) | |
| model.cuda() | |
| model.eval() | |
| print(f" Model încărcat în {time.time() - t_load:.1f}s") | |
| # Latente vocale (cache) | |
| if os.path.exists(CACHED_LATENTS): | |
| print(f"\nÎncărcare latente din cache...") | |
| cached = torch.load(CACHED_LATENTS, weights_only=True) | |
| gpt_cond_latent = cached["gpt_cond_latent"].cuda() | |
| speaker_embedding = cached["speaker_embedding"].cuda() | |
| print(" Latente încărcate din cache!") | |
| else: | |
| print(f"\nExtragere latente din: {os.path.basename(REFERENCE_AUDIO)}") | |
| gpt_cond_latent, speaker_embedding = model.get_conditioning_latents( | |
| audio_path=[REFERENCE_AUDIO], | |
| gpt_cond_len=GPT_COND_LEN, | |
| ) | |
| # Salvare cache | |
| torch.save({ | |
| "gpt_cond_latent": gpt_cond_latent.cpu(), | |
| "speaker_embedding": speaker_embedding.cpu(), | |
| }, CACHED_LATENTS) | |
| print(f" Latente salvate în cache: {CACHED_LATENTS}") | |
| # Pregătire text | |
| text = TEXT.strip() | |
| if not text: | |
| print("ERROR: TEXT este gol!") | |
| sys.exit(1) | |
| text_clean = text.translate(CEDILLA_TO_COMMA) | |
| # Generare | |
| print(f"\nGenerare audio...") | |
| print(f" Text: {text_clean[:100]}{'...' if len(text_clean) > 100 else ''}") | |
| t_gen = time.time() | |
| with torch.no_grad(): | |
| output = model.inference( | |
| text_clean, | |
| "ro", | |
| gpt_cond_latent, | |
| speaker_embedding, | |
| temperature=TEMPERATURE, | |
| top_p=TOP_P, | |
| top_k=TOP_K, | |
| length_penalty=LENGTH_PENALTY, | |
| repetition_penalty=REPETITION_PENALTY, | |
| ) | |
| wav = output["wav"] | |
| if isinstance(wav, torch.Tensor): | |
| wav = wav.cpu().numpy() | |
| wav = wav.squeeze() | |
| gen_time = time.time() - t_gen | |
| duration = len(wav) / 24000 | |
| # Salvare | |
| from datetime import datetime | |
| timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") | |
| out_path = os.path.join(OUTPUT_DIR, f"tts_{timestamp}.wav") | |
| sf.write(out_path, wav, 24000) | |
| # Raport | |
| print(f"\n{'=' * 50}") | |
| print(f" Audio generat: {out_path}") | |
| print(f" Durată audio: {duration:.2f}s") | |
| print(f" Timp generare: {gen_time:.2f}s") | |
| print(f" Timp total: {time.time() - t_start:.1f}s") | |
| print(f" RTF: {gen_time / duration:.2f}x") | |
| print(f"{'=' * 50}") | |
| if __name__ == "__main__": | |
| main() | |