xtts-v2-moldovan-romanian / generate_tts.py
FraPiz's picture
Publish Moldovan Romanian XTTS-v2 model
209ef39 verified
Raw History Blame Contribute Delete
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()