""" 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()