nc-murray/spectrogram-reconstruction
0
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 