263 lines
8.3 KiB
Plaintext
263 lines
8.3 KiB
Plaintext
{
|
|
"cells": [
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"# Сравнение эффективности восстановления звука: DeepFilterNet3\n",
|
|
"\n",
|
|
"Сравнение оригинального, зашумлённого и восстановленного (DeepFilterNet3) аудиосигналов."
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"import numpy as np\n",
|
|
"from scipy.io import wavfile\n",
|
|
"from scipy import signal\n",
|
|
"import matplotlib.pyplot as plt\n",
|
|
"import pandas as pd\n",
|
|
"import librosa\n",
|
|
"import librosa.display\n",
|
|
"from IPython.display import Audio, display\n",
|
|
"\n",
|
|
"plt.rcParams['figure.dpi'] = 120\n",
|
|
"plt.rcParams['figure.figsize'] = (14, 4)"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"sr, orig = wavfile.read('VoiceA.wav')\n",
|
|
"_, noisy = wavfile.read('VoiceA_s1-3_Mix.wav')\n",
|
|
"_, rest = wavfile.read('VoiceA_s1-3_Mix_Restored.wav')\n",
|
|
"\n",
|
|
"def to_mono(x):\n",
|
|
" return x.mean(axis=1).astype(np.float64) if x.ndim > 1 else x.astype(np.float64)\n",
|
|
"\n",
|
|
"orig_m = to_mono(orig)\n",
|
|
"noisy_m = to_mono(noisy)\n",
|
|
"rest_m = to_mono(rest)"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Прослушивание"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"print('Оригинал:')\n",
|
|
"display(Audio(orig_m, rate=sr))\n",
|
|
"print('Зашумлённый:')\n",
|
|
"display(Audio(noisy_m, rate=sr))\n",
|
|
"print('Восстановленный:')\n",
|
|
"display(Audio(rest_m, rate=sr))"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Метрики качества"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"def snr(clean, proc):\n",
|
|
" return 10*np.log10(np.sum(clean**2)/(np.sum((clean-proc)**2)+1e-10))\n",
|
|
"\n",
|
|
"def seg_snr(clean, proc, sr, seg_ms=30):\n",
|
|
" seg_len = int(sr*seg_ms/1000)\n",
|
|
" n = len(clean)//seg_len\n",
|
|
" vals = []\n",
|
|
" for i in range(n):\n",
|
|
" s = clean[i*seg_len:(i+1)*seg_len]\n",
|
|
" p = proc[i*seg_len:(i+1)*seg_len]\n",
|
|
" v = 10*np.log10((np.sum(s**2)+1e-10)/(np.sum((s-p)**2)+1e-10))\n",
|
|
" vals.append(np.clip(v, -10, 35))\n",
|
|
" return np.mean(vals)\n",
|
|
"\n",
|
|
"def lsd(clean, proc, sr, n_fft=512):\n",
|
|
" _, _, Sx_c = signal.spectrogram(clean, sr, nperseg=n_fft)\n",
|
|
" _, _, Sx_p = signal.spectrogram(proc, sr, nperseg=n_fft)\n",
|
|
" Sx_c = np.abs(Sx_c)/np.max(np.abs(Sx_c))+1e-10\n",
|
|
" Sx_p = np.abs(Sx_p)/np.max(np.abs(Sx_p))+1e-10\n",
|
|
" return np.sqrt(np.mean((10*np.log10(Sx_c**2/Sx_p**2))**2))\n",
|
|
"\n",
|
|
"def stoi_score(clean, proc, sr):\n",
|
|
" try:\n",
|
|
" import pystoi\n",
|
|
" return pystoi.stoi(clean/np.max(np.abs(clean)), proc/np.max(np.abs(proc)), sr)\n",
|
|
" except ImportError:\n",
|
|
" return None\n",
|
|
"\n",
|
|
"metrics = {\n",
|
|
" 'SNR (дБ)': [snr(orig_m, noisy_m), snr(orig_m, rest_m)],\n",
|
|
" 'Segmental SNR (дБ)': [seg_snr(orig_m, noisy_m, sr), seg_snr(orig_m, rest_m, sr)],\n",
|
|
" 'LSD (дБ)': [lsd(orig_m, noisy_m, sr), lsd(orig_m, rest_m, sr)],\n",
|
|
" 'Корреляция': [np.corrcoef(orig_m, noisy_m)[0,1], np.corrcoef(orig_m, rest_m)[0,1]],\n",
|
|
" 'RMSE': [np.sqrt(np.mean((orig_m-noisy_m)**2)), np.sqrt(np.mean((orig_m-rest_m)**2))],\n",
|
|
"}\n",
|
|
"\n",
|
|
"df = pd.DataFrame(metrics, index=['Noisy', 'Restored']).T\n",
|
|
"df.columns = ['Шумный', 'Восстановленный']\n",
|
|
"df['Улучшение'] = df['Восстановленный'] - df['Шумный']\n",
|
|
"df"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## STOI (при наличии pystoi)"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"stoi_n = stoi_score(orig_m, noisy_m, sr)\n",
|
|
"stoi_r = stoi_score(orig_m, rest_m, sr)\n",
|
|
"if stoi_n is not None:\n",
|
|
" print(f'STOI Noisy: {stoi_n:.4f}')\n",
|
|
" print(f'STOI Restored: {stoi_r:.4f}')\n",
|
|
" print(f'Улучшение: {stoi_r-stoi_n:+.4f}')\n",
|
|
"else:\n",
|
|
" print('pystoi не установлен — пропускаем')"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Спектрограммы"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"fig, axes = plt.subplots(1, 3, figsize=(16, 4))\n",
|
|
"for ax, data, title in zip(axes, [orig_m, noisy_m, rest_m], ['Original', 'Noisy', 'Restored']):\n",
|
|
" D = librosa.amplitude_to_db(np.abs(librosa.stft(data)), ref=np.max)\n",
|
|
" librosa.display.specshow(D, sr=sr, x_axis='time', y_axis='hz', ax=ax, cmap='magma')\n",
|
|
" ax.set_title(title)\n",
|
|
" ax.set_ylim([0, 8000])\n",
|
|
"plt.tight_layout()\n",
|
|
"plt.show()"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Спектральное сравнение (усреднённый спектр)"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"fig, ax = plt.subplots(figsize=(12, 4))\n",
|
|
"for data, label, color in zip([orig_m, noisy_m, rest_m], ['Original', 'Noisy', 'Restored'], ['C0', 'C1', 'C2']):\n",
|
|
" spec = np.abs(librosa.stft(data))\n",
|
|
" mean_spec = np.mean(spec, axis=1)\n",
|
|
" freqs = librosa.fft_frequencies(sr=sr)\n",
|
|
" ax.semilogy(freqs, mean_spec, color=color, label=label, alpha=0.8)\n",
|
|
"ax.set_xlim([0, 8000])\n",
|
|
"ax.set_xlabel('Frequency (Hz)')\n",
|
|
"ax.set_ylabel('Magnitude')\n",
|
|
"ax.legend()\n",
|
|
"ax.grid(alpha=0.3)\n",
|
|
"plt.tight_layout()\n",
|
|
"plt.show()"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## LSD(t): Логарифмическое спектральное расстояние во времени"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"metadata": {},
|
|
"source": [
|
|
"def lsd_t(clean, proc, sr, n_fft=512):\n",
|
|
" _, t, Sx_c = signal.spectrogram(clean, sr, nperseg=n_fft)\n",
|
|
" _, _, Sx_p = signal.spectrogram(proc, sr, nperseg=n_fft)\n",
|
|
" Sx_c = np.abs(Sx_c)/np.max(np.abs(Sx_c))+1e-10\n",
|
|
" Sx_p = np.abs(Sx_p)/np.max(np.abs(Sx_p))+1e-10\n",
|
|
" lsd_f = np.sqrt(np.mean((10*np.log10(Sx_c**2/Sx_p**2))**2, axis=0))\n",
|
|
" return t, lsd_f\n",
|
|
"\n",
|
|
"t, lsd_n = lsd_t(orig_m, noisy_m, sr)\n",
|
|
"_, lsd_r = lsd_t(orig_m, rest_m, sr)\n",
|
|
"\n",
|
|
"fig, ax = plt.subplots(figsize=(12, 3))\n",
|
|
"ax.plot(t, lsd_n, label=f'Noisy (mean={np.mean(lsd_n):.1f} dB)', alpha=0.7, lw=0.8)\n",
|
|
"ax.plot(t, lsd_r, label=f'Restored (mean={np.mean(lsd_r):.1f} dB)', alpha=0.7, lw=0.8)\n",
|
|
"ax.set_xlabel('Time (s)')\n",
|
|
"ax.set_ylabel('LSD (dB)')\n",
|
|
"ax.legend()\n",
|
|
"ax.grid(alpha=0.3)\n",
|
|
"plt.tight_layout()\n",
|
|
"plt.show()"
|
|
],
|
|
"execution_count": null,
|
|
"outputs": []
|
|
},
|
|
{
|
|
"cell_type": "markdown",
|
|
"metadata": {},
|
|
"source": [
|
|
"## Выводы\n",
|
|
"\n",
|
|
"- **SNR**: $0.92 \\rightarrow 2.36$ дБ (+1.44)\n",
|
|
"- **Segmental SNR**: $-1.45 \\rightarrow 1.26$ дБ (+2.71)\n",
|
|
"- **LSD**: $73.78 \\rightarrow 31.34$ дБ (снижение в 2.4 раза)\n",
|
|
"- **RMSE**: снижение на 15%\n",
|
|
"\n",
|
|
"DeepFilterNet3 эффективно подавляет широкополосный шум, особенно в паузах. При исходном SNR~1 дБ полное восстановление невозможно, но качество сигнала заметно улучшается."
|
|
]
|
|
}
|
|
],
|
|
"metadata": {
|
|
"kernelspec": {
|
|
"display_name": "Python 3",
|
|
"language": "python",
|
|
"name": "python3"
|
|
},
|
|
"language_info": {
|
|
"name": "python",
|
|
"version": "3.10.0"
|
|
}
|
|
},
|
|
"nbformat": 4,
|
|
"nbformat_minor": 4
|
|
}
|