Team Ai
Apppublic

nc-murray/spectrogram-reconstruction

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
visualizer.py153 linesDownload Raw Back to specrec
1"""Reconstruction comparison plots and accuracy curves."""2 3import numpy as np4import matplotlib5import matplotlib.pyplot as plt6import matplotlib.gridspec as gridspec7import librosa8import librosa.display9 10 11def plot_reconstruction_comparison(12    original: np.ndarray,13    reconstructed: np.ndarray,14    sr: int,15    title: str = "Reconstruction Comparison",16    output_path: str = None,17    n_fft: int = 2048,18    hop_length: int = 512,19) -> plt.Figure:20    """21    2×2 grid:22      top-left:  original waveform23      top-right: reconstructed waveform24      bot-left:  original spectrogram25      bot-right: reconstructed spectrogram26    """27    # Trim/pad so both are the same length for visual alignment28    n = min(len(original), len(reconstructed))29    orig = original[:n]30    recon = reconstructed[:n]31 32    S_orig  = librosa.amplitude_to_db(33        np.abs(librosa.stft(orig,  n_fft=n_fft, hop_length=hop_length)), ref=np.max)34    S_recon = librosa.amplitude_to_db(35        np.abs(librosa.stft(recon, n_fft=n_fft, hop_length=hop_length)), ref=np.max)36 37    t = np.linspace(0, n / sr, n)38 39    fig = plt.figure(figsize=(14, 7))40    fig.suptitle(title, fontsize=13, fontweight="bold")41    gs = gridspec.GridSpec(2, 2, figure=fig, hspace=0.45, wspace=0.35)42 43    vmin = max(S_orig.min(), S_recon.min(), -80)44 45    # Waveforms46    for col, (label, wave) in enumerate([("Original", orig), ("Reconstructed", recon)]):47        ax = fig.add_subplot(gs[0, col])48        ax.plot(t, wave, linewidth=0.6, color="steelblue" if col == 0 else "darkorange")49        ax.set_title(f"{label} waveform", fontsize=10)50        ax.set_xlabel("Time (s)")51        ax.set_ylabel("Amplitude")52        ax.set_xlim(t[0], t[-1])53        ax.grid(True, linewidth=0.3, alpha=0.5)54 55    # Spectrograms56    for col, (label, S) in enumerate([("Original", S_orig), ("Reconstructed", S_recon)]):57        ax = fig.add_subplot(gs[1, col])58        img = librosa.display.specshow(59            S, sr=sr, hop_length=hop_length, x_axis="time", y_axis="hz",60            ax=ax, cmap="viridis", vmin=vmin, vmax=0,61        )62        ax.set_title(f"{label} spectrogram", fontsize=10)63        ax.set_xlabel("Time (s)")64        ax.set_ylabel("Frequency (Hz)")65        fig.colorbar(img, ax=ax, format="%+2.0f dB", pad=0.02)66 67    if output_path:68        fig.savefig(output_path, dpi=120, bbox_inches="tight")69        plt.close(fig)70    return fig71 72 73def plot_accuracy_vs_iterations(74    results: dict,75    title: str = "Reconstruction Accuracy vs Griffin-Lim Iterations",76    output_path: str = None,77) -> plt.Figure:78    """79    Dual-axis line plot: spectral convergence (left) and SNR dB (right)80    versus n_iter.  results is the dict returned by run_roundtrip_accuracy_test().81    """82    iters = sorted(results.keys())83    sc_values  = [results[i]["spectral_convergence"] for i in iters]84    snr_values = [results[i]["snr_db"]               for i in iters]85 86    fig, ax1 = plt.subplots(figsize=(8, 4.5))87    fig.suptitle(title, fontsize=12, fontweight="bold")88 89    color_sc  = "steelblue"90    color_snr = "darkorange"91 92    ax1.plot(iters, sc_values, "o-", color=color_sc, linewidth=2, markersize=6, label="Spectral convergence")93    ax1.set_xlabel("Griffin-Lim iterations")94    ax1.set_ylabel("Spectral convergence  (↓ better)", color=color_sc)95    ax1.tick_params(axis="y", labelcolor=color_sc)96    ax1.set_xticks(iters)97    ax1.grid(True, linewidth=0.3, alpha=0.5)98 99    ax2 = ax1.twinx()100    ax2.plot(iters, snr_values, "s--", color=color_snr, linewidth=2, markersize=6, label="SNR (dB)")101    ax2.set_ylabel("SNR  (dB)  (↑ better)", color=color_snr)102    ax2.tick_params(axis="y", labelcolor=color_snr)103 104    # Combined legend105    lines1, labels1 = ax1.get_legend_handles_labels()106    lines2, labels2 = ax2.get_legend_handles_labels()107    ax1.legend(lines1 + lines2, labels1 + labels2, loc="center right", fontsize=9)108 109    fig.tight_layout()110    if output_path:111        fig.savefig(output_path, dpi=120, bbox_inches="tight")112        plt.close(fig)113    return fig114 115 116def plot_colormap_sensitivity(117    cmap_results: dict,118    title: str = "Colormap Sensitivity",119    output_path: str = None,120) -> plt.Figure:121    """122    Bar chart: spectral convergence per colormap.123    cmap_results is the dict returned by run_colormap_sensitivity_test().124    """125    cmaps = list(cmap_results.keys())126    sc_values  = [cmap_results[c]["spectral_convergence"] for c in cmaps]127    snr_values = [cmap_results[c]["snr_db"]               for c in cmaps]128 129    x = np.arange(len(cmaps))130    width = 0.35131 132    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10, 4))133    fig.suptitle(title, fontsize=12, fontweight="bold")134 135    ax1.bar(x, sc_values, width=0.6, color="steelblue", alpha=0.85)136    ax1.set_xticks(x); ax1.set_xticklabels(cmaps)137    ax1.set_ylabel("Spectral convergence  (↓ better)")138    ax1.set_title("Spectral convergence by colormap")139    ax1.grid(axis="y", linewidth=0.3, alpha=0.5)140    ax1.set_ylim(0, max(sc_values) * 1.2)141 142    ax2.bar(x, snr_values, width=0.6, color="darkorange", alpha=0.85)143    ax2.set_xticks(x); ax2.set_xticklabels(cmaps)144    ax2.set_ylabel("SNR (dB)  (↑ better)")145    ax2.set_title("SNR by colormap")146    ax2.grid(axis="y", linewidth=0.3, alpha=0.5)147 148    fig.tight_layout()149    if output_path:150        fig.savefig(output_path, dpi=120, bbox_inches="tight")151        plt.close(fig)152    return fig153