| 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 |
|
|
| |
| 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] |
|
|
| |
| model_configs = { |
| |
| |
| "brais_final_slm": "Configs/inference_config.yml", |
| |
| } |
|
|
|
|
| |
| 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'): |
| |
| 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) |
| |
| 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 |
| |
| 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() |
| |
| 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") |
|
|