Nos_StyleTTS2-Brais-GL / scripts /grid_search_inference.py
cmagui's picture
Initial commit: full repository with code, configs and weights
613ce86
Raw
History Blame Contribute Delete
7.13 kB
import subprocess
import itertools
import os
import re
import matplotlib.pyplot as plt
import numpy as np
import torchaudio
import librosa
from speechmos import dnsmos
# Parámetros a probar
diffusion_steps = [5, 10]
embedding_scales = [1.0, 1.5]
alphas = [0, 0.2, 0.4, 0.6, 0.8, 1]
betas = [0, 0.2, 0.4, 0.6, 0.8, 1]
ts = [0, 0.2, 0.4, 0.6, 0.8, 1]
# Modelos/configuraciones a probar
model_configs = {
# "brais_no_slm_oversampling": "Configs/inference_config_brais.yml",
# "celtia_slm_acentos": "Configs/inference_config_celtia.yml",
"brais_final_slm": "Configs/inference_config.yml",
# Agrega más rutas si tienes más configuraciones/modelos
}
# Textos a sintetizar
files = [
"Demo/pasaxe.txt",
]
output_base = "outputs/grid_search"
os.makedirs(output_base, exist_ok=True)
def compute_dnsmos_from_wav(wav_path, key='ovrl_mos'):
# Carga el audio, resamplea a 16kHz y calcula la métrica
wav, sr = torchaudio.load(wav_path)
wav = wav.squeeze().numpy()
if sr != 16000:
wav = librosa.resample(
wav, orig_sr=sr, target_sr=16000, res_type='kaiser_best', fix=True)
if wav is not None and len(wav) > 0:
mos_dict = dnsmos.run(wav, sr=16000)
val = mos_dict.get(key, None)
if val is not None:
try:
return float(val)
except Exception:
return None
return None
best_overall = {
'dns_mos': -np.inf,
'config': None
}
for model_name, model_config in model_configs.items():
for file in files:
for t in ts:
best = {
'dns_mos': -np.inf,
'config': None
}
results = np.zeros((len(alphas), len(betas)))
for diffusion_step in diffusion_steps:
for embedding_scale in embedding_scales:
for i, alpha in enumerate(alphas):
for j, beta in enumerate(betas):
out_dir = os.path.join(
output_base, model_name, f"d{diffusion_step}_e{embedding_scale}")
out_file = f"{model_name}_a{alpha}_b{beta}_t{t}"
os.makedirs(out_dir, exist_ok=True)
cmd = [
"python3", "inference.py",
"--config", model_config,
"--file", file,
"--device", "0",
"--output_dir", out_dir,
"--output_file", out_file,
"--evaluate",
"--alpha", str(alpha),
"--beta", str(beta),
"--t", str(t),
"--diffusion_steps", str(diffusion_step),
"--embedding_scale", str(embedding_scale)
]
print(f"Ejecutando: {cmd} en {out_dir}")
proc = subprocess.run(
cmd, capture_output=True, text=True)
# Buscar el archivo de audio generado
wav_file = os.path.join(
out_dir, f"{out_file}_{diffusion_step}_{embedding_scale}.wav")
if os.path.exists(wav_file):
dns_mos = compute_dnsmos_from_wav(
wav_file, key='ovrl_mos')
print(
f"DNSMOS obtenido: {dns_mos} para alpha={alpha}, beta={beta}, t={t}, diffusion_step={diffusion_step}, embedding_scale={embedding_scale}")
if dns_mos is not None:
results[i, j] = dns_mos
if dns_mos > best['dns_mos']:
best['dns_mos'] = dns_mos
best['config'] = {
'model_name': model_name,
'model_config': model_config,
'file': file,
't': t,
'alpha': alpha,
'beta': beta,
'diffusion_step': diffusion_step,
'embedding_scale': embedding_scale
}
if dns_mos > best_overall['dns_mos']:
best_overall['dns_mos'] = dns_mos
best_overall['config'] = {
'model_name': model_name,
'model_config': model_config,
'file': file,
't': t,
'alpha': alpha,
'beta': beta,
'diffusion_step': diffusion_step,
'embedding_scale': embedding_scale
}
else:
results[i, j] = np.nan
else:
print(
f"No se encontró el archivo de audio: {wav_file}")
results[i, j] = np.nan
# Graficar mapa 2D para este t
plt.figure(figsize=(8, 6))
plt.imshow(results, origin='lower', aspect='auto',
extent=[min(betas), max(betas),
min(alphas), max(alphas)],
cmap='viridis')
plt.colorbar(label='DNSMOS')
plt.xlabel('Beta')
plt.ylabel('Alpha')
plt.title(
f'DNSMOS para t={t}, modelo={model_name}, d={diffusion_step}, e={embedding_scale}')
plt.xticks(betas)
plt.yticks(alphas)
plt.tight_layout()
plot_path = os.path.join(
output_base, f"dnsmap_{model_name}_t{t}_d{diffusion_step}_e{embedding_scale}.png")
plt.savefig(plot_path)
plt.close()
# Mostrar mejor resultado para este t
print(
f"\nMejor DNSMOS para modelo={model_name}, t={t}: {best['dns_mos']}\nConfiguración: {best['config']}\n")
print("\n==============================")
print("Mejor DNSMOS global:")
print(f"DNSMOS: {best_overall['dns_mos']}")
print(f"Configuración: {best_overall['config']}")
print("==============================\n")