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