JavohirTF7 hshr commited on
Commit
5e0f239
·
0 Parent(s):

Duplicate from hshr/DeepFilterNet2

Browse files

Co-authored-by: Hendrik Schröter <hshr@users.noreply.huggingface.co>

.flake8 ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [flake8]
2
+ ignore = E203, E266, E501, W503
3
+ max-line-length = 100
4
+ import-order-style = google
5
+ application-import-names = flake8
6
+ select = B,C,E,F,W,T4,B9
7
+ exclude =
8
+ .tox,
9
+ .git,
10
+ __pycache__,
11
+ docs,
12
+ sbatch,
13
+ .venv,
14
+ *.pyc,
15
+ *.egg-info,
16
+ .cache,
17
+ .eggs
.gitattributes ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ftz filter=lfs diff=lfs merge=lfs -text
6
+ *.gz filter=lfs diff=lfs merge=lfs -text
7
+ *.h5 filter=lfs diff=lfs merge=lfs -text
8
+ *.joblib filter=lfs diff=lfs merge=lfs -text
9
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
10
+ *.model filter=lfs diff=lfs merge=lfs -text
11
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
12
+ *.onnx filter=lfs diff=lfs merge=lfs -text
13
+ *.ot filter=lfs diff=lfs merge=lfs -text
14
+ *.parquet filter=lfs diff=lfs merge=lfs -text
15
+ *.pb filter=lfs diff=lfs merge=lfs -text
16
+ *.pt filter=lfs diff=lfs merge=lfs -text
17
+ *.pth filter=lfs diff=lfs merge=lfs -text
18
+ *.rar filter=lfs diff=lfs merge=lfs -text
19
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
20
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
21
+ *.tflite filter=lfs diff=lfs merge=lfs -text
22
+ *.tgz filter=lfs diff=lfs merge=lfs -text
23
+ *.wasm filter=lfs diff=lfs merge=lfs -text
24
+ *.xz filter=lfs diff=lfs merge=lfs -text
25
+ *.zip filter=lfs diff=lfs merge=lfs -text
26
+ *.zstandard filter=lfs diff=lfs merge=lfs -text
27
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
28
+ *ckpt filter=lfs diff=lfs merge=lfs -text
29
+ *best filter=lfs diff=lfs merge=lfs -text
30
+ *wav filter=lfs diff=lfs merge=lfs -text
.gitignore ADDED
@@ -0,0 +1,145 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Own stuff
2
+ *.wav
3
+ *.png
4
+ *.pdf
5
+ out/
6
+ export/
7
+ DeepFilterNet/poetry.lock
8
+ gradio_cached_examples/
9
+
10
+ ### Rust gitignore ###
11
+
12
+ # Generated by Cargo
13
+ # will have compiled files and executables
14
+ debug/
15
+ target/
16
+
17
+ # Remove Cargo.lock from gitignore if creating an executable, leave it for libraries
18
+ # More information here https://doc.rust-lang.org/cargo/guide/cargo-toml-vs-cargo-lock.html
19
+ Cargo.lock
20
+
21
+ # These are backup files generated by rustfmt
22
+ **/*.rs.bk
23
+
24
+ ### Python gitignore ###
25
+
26
+ # Byte-compiled / optimized / DLL files
27
+ __pycache__/
28
+ *.py[cod]
29
+ *$py.class
30
+
31
+ # C extensions
32
+ *.so
33
+
34
+ # Distribution / packaging
35
+ .Python
36
+ build/
37
+ develop-eggs/
38
+ dist/
39
+ downloads/
40
+ eggs/
41
+ .eggs/
42
+ lib/
43
+ lib64/
44
+ parts/
45
+ sdist/
46
+ var/
47
+ wheels/
48
+ pip-wheel-metadata/
49
+ share/python-wheels/
50
+ *.egg-info/
51
+ .installed.cfg
52
+ *.egg
53
+ MANIFEST
54
+
55
+ # PyInstaller
56
+ # Usually these files are written by a python script from a template
57
+ # before PyInstaller builds the exe, so as to inject date/other infos into it.
58
+ *.manifest
59
+ *.spec
60
+
61
+ # Installer logs
62
+ pip-log.txt
63
+ pip-delete-this-directory.txt
64
+
65
+ # Unit test / coverage reports
66
+ typings
67
+ htmlcov/
68
+ .tox/
69
+ .nox/
70
+ .coverage
71
+ .coverage.*
72
+ .cache
73
+ nosetests.xml
74
+ coverage.xml
75
+ *.cover
76
+ .hypothesis/
77
+ .pytest_cache/
78
+
79
+ # Translations
80
+ *.mo
81
+ *.pot
82
+
83
+ # Django stuff:
84
+ *.log
85
+ local_settings.py
86
+ db.sqlite3
87
+
88
+ # Flask stuff:
89
+ instance/
90
+ .webassets-cache
91
+
92
+ # Scrapy stuff:
93
+ .scrapy
94
+
95
+ # Sphinx documentation
96
+ docs/_build/
97
+
98
+ # PyBuilder
99
+ target/
100
+
101
+ # Jupyter Notebook
102
+ .ipynb_checkpoints
103
+
104
+ # IPython
105
+ profile_default/
106
+ ipython_config.py
107
+
108
+ # pyenv
109
+ .python-version
110
+
111
+ # celery beat schedule file
112
+ celerybeat-schedule
113
+
114
+ # SageMath parsed files
115
+ *.sage.py
116
+
117
+ # Environments
118
+ .env
119
+ .venv
120
+ env/
121
+ venv/
122
+ ENV/
123
+ env.bak/
124
+ venv.bak/
125
+
126
+ # Spyder project settings
127
+ .spyderproject
128
+ .spyproject
129
+
130
+ # Rope project settings
131
+ .ropeproject
132
+
133
+ # mkdocs documentation
134
+ /site
135
+
136
+ # mypy
137
+ .mypy_cache/
138
+ .dmypy.json
139
+ dmypy.json
140
+
141
+ # Pyre type checker
142
+ .pyre/
143
+
144
+ # IDE
145
+ .idea
DeepFilterNet2/checkpoints/model_96.ckpt.best ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:bb5eccb429e675bb4ec5ec9e280f048bfff9787b40bd3eb835fd11509eb14a3e
3
+ size 9397209
DeepFilterNet2/config.ini ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [train]
2
+ seed = 43
3
+ device =
4
+ model = deepfilternet2
5
+ jit = false
6
+ mask_only = false
7
+ df_only = false
8
+ batch_size = 96
9
+ batch_size_eval = 128
10
+ num_workers = 16
11
+ max_sample_len_s = 3.0
12
+ p_atten_lim = 0.0
13
+ p_reverb = 0.1
14
+ overfit = false
15
+ max_epochs = 100
16
+ log_freq = 100
17
+ log_timings = False
18
+ validation_criteria = loss
19
+ validation_criteria_rule = min
20
+ early_stopping_patience = 15
21
+ global_ds_sampling_f = 1
22
+ num_prefetch_batches = 8
23
+ dataloader_snrs = -5,0,5,10,20,40
24
+ detect_anomaly = false
25
+ batch_size_scheduling = 0/8,1/16,2/24,5/32,10/64,20/128,40/9999
26
+ start_eval = true
27
+ validation_set_caching = false
28
+
29
+ [df]
30
+ sr = 48000
31
+ fft_size = 960
32
+ hop_size = 480
33
+ nb_erb = 32
34
+ nb_df = 96
35
+ norm_tau = 1
36
+ lsnr_max = 35
37
+ lsnr_min = -15
38
+ min_nb_erb_freqs = 2
39
+ pad_mode = input_specf
40
+
41
+ [deepfilternet]
42
+ conv_lookahead = 2
43
+ conv_ch = 64
44
+ conv_depthwise = True
45
+ emb_hidden_dim = 256
46
+ emb_num_layers = 3
47
+ gru_groups = 8
48
+ linear_groups = 8
49
+ conv_dec_mode = transposed
50
+ convt_depthwise = True
51
+ mask_pf = False
52
+ df_order = 5
53
+ df_lookahead = 2
54
+ df_hidden_dim = 256
55
+ df_num_layers = 2
56
+ dfop_method = df
57
+ group_shuffle = False
58
+ conv_kernel = 1,3
59
+ df_gru_skip = none
60
+ df_output_layer = groupedlinear
61
+ gru_type = squeeze
62
+ df_pathway_kernel_size_t = 5
63
+ df_n_iter = 1
64
+ enc_concat = True
65
+ conv_kernel_inp = 3,3
66
+
67
+ [localsnrloss]
68
+ factor = 1e-3
69
+
70
+ [maskloss]
71
+ factor = 0
72
+ mask = iam
73
+ gamma = 0.6
74
+ gamma_pred = 0.6
75
+ f_under = 1
76
+
77
+ [spectralloss]
78
+ factor_magnitude = 1000
79
+ factor_complex = 1000
80
+ gamma = 0.3
81
+
82
+ [dfalphaloss]
83
+ factor = 0.0
84
+
85
+ [multiresspecloss]
86
+ factor = 500
87
+ factor_complex = 500
88
+ gamma = 0.3
89
+ fft_sizes = 256,512,1024
90
+
91
+ [optim]
92
+ lr = 0.001
93
+ momentum = 0
94
+ weight_decay = 1e-12
95
+ weight_decay_end = 0.05
96
+ optimizer = adamw
97
+ lr_min = 1e-06
98
+ lr_warmup = 0.0001
99
+ warmup_epochs = 3
100
+ lr_cycle_mul = 1.0
101
+ lr_cycle_decay = 0.5
102
+ lr_cycle_limit = 1
103
+ lr_update_per_epoch = False
104
+ lr_cycle_epochs = -1
105
+
106
+ [sdrloss]
107
+ factor = 0.0
108
+ segmental_ws = 0
109
+
README.md ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ title: DeepFilterNet2
3
+ emoji: 💩
4
+ colorFrom: gray
5
+ colorTo: red
6
+ sdk: gradio
7
+ app_file: app.py
8
+ sdk_version: 3.17.1
9
+ pinned: false
10
+ license: apache-2.0
11
+ duplicated_from: hshr/DeepFilterNet2
12
+ ---
13
+
14
+ Check out the configuration reference at https://huggingface.co/docs/hub/spaces#reference
app.py ADDED
@@ -0,0 +1,313 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ import tempfile
3
+ from typing import Optional, Tuple, Union
4
+
5
+ import gradio as gr
6
+ import matplotlib.pyplot as plt
7
+ import numpy as np
8
+ import torch
9
+ from loguru import logger
10
+ from PIL import Image
11
+ from torch import Tensor
12
+ from torchaudio.backend.common import AudioMetaData
13
+
14
+ from df import config
15
+ from df.enhance import enhance, init_df, load_audio, save_audio
16
+ from df.io import resample
17
+
18
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
19
+ model, df, _ = init_df("./DeepFilterNet2", config_allow_defaults=True)
20
+ model = model.to(device=device).eval()
21
+
22
+ fig_noisy: plt.Figure
23
+ fig_enh: plt.Figure
24
+ ax_noisy: plt.Axes
25
+ ax_enh: plt.Axes
26
+ fig_noisy, ax_noisy = plt.subplots(figsize=(15.2, 4))
27
+ fig_noisy.set_tight_layout(True)
28
+ fig_enh, ax_enh = plt.subplots(figsize=(15.2, 4))
29
+ fig_enh.set_tight_layout(True)
30
+
31
+ NOISES = {
32
+ "None": None,
33
+ "Kitchen": "samples/dkitchen.wav",
34
+ "Living Room": "samples/dliving.wav",
35
+ "River": "samples/nriver.wav",
36
+ "Cafe": "samples/scafe.wav",
37
+ }
38
+
39
+
40
+ def mix_at_snr(clean, noise, snr, eps=1e-10):
41
+ """Mix clean and noise signal at a given SNR.
42
+
43
+ Args:
44
+ clean: 1D Tensor with the clean signal to mix.
45
+ noise: 1D Tensor of shape.
46
+ snr: Signal to noise ratio.
47
+
48
+ Returns:
49
+ clean: 1D Tensor with gain changed according to the snr.
50
+ noise: 1D Tensor with the combined noise channels.
51
+ mix: 1D Tensor with added clean and noise signals.
52
+
53
+ """
54
+ clean = torch.as_tensor(clean).mean(0, keepdim=True)
55
+ noise = torch.as_tensor(noise).mean(0, keepdim=True)
56
+ if noise.shape[1] < clean.shape[1]:
57
+ noise = noise.repeat((1, int(math.ceil(clean.shape[1] / noise.shape[1]))))
58
+ max_start = int(noise.shape[1] - clean.shape[1])
59
+ start = torch.randint(0, max_start, ()).item() if max_start > 0 else 0
60
+ logger.debug(f"start: {start}, {clean.shape}")
61
+ noise = noise[:, start : start + clean.shape[1]]
62
+ E_speech = torch.mean(clean.pow(2)) + eps
63
+ E_noise = torch.mean(noise.pow(2))
64
+ K = torch.sqrt((E_noise / E_speech) * 10 ** (snr / 10) + eps)
65
+ noise = noise / K
66
+ mixture = clean + noise
67
+ logger.debug("mixture: {mixture.shape}")
68
+ assert torch.isfinite(mixture).all()
69
+ max_m = mixture.abs().max()
70
+ if max_m > 1:
71
+ logger.warning(f"Clipping detected during mixing. Reducing gain by {1/max_m}")
72
+ clean, noise, mixture = clean / max_m, noise / max_m, mixture / max_m
73
+ return clean, noise, mixture
74
+
75
+
76
+ def load_audio_gradio(
77
+ audio_or_file: Union[None, str, Tuple[int, np.ndarray]], sr: int
78
+ ) -> Optional[Tuple[Tensor, AudioMetaData]]:
79
+ if audio_or_file is None:
80
+ return None
81
+ if isinstance(audio_or_file, str):
82
+ if audio_or_file.lower() == "none":
83
+ return None
84
+ # First try default format
85
+ audio, meta = load_audio(audio_or_file, sr)
86
+ else:
87
+ meta = AudioMetaData(-1, -1, -1, -1, "")
88
+ assert isinstance(audio_or_file, (tuple, list))
89
+ meta.sample_rate, audio_np = audio_or_file
90
+ # Gradio documentation says, the shape is [samples, 2], but apparently sometimes its not.
91
+ audio_np = audio_np.reshape(audio_np.shape[0], -1).T
92
+ if audio_np.dtype == np.int16:
93
+ audio_np = (audio_np / (1 << 15)).astype(np.float32)
94
+ elif audio_np.dtype == np.int32:
95
+ audio_np = (audio_np / (1 << 31)).astype(np.float32)
96
+ audio = resample(torch.from_numpy(audio_np), meta.sample_rate, sr)
97
+ return audio, meta
98
+
99
+
100
+ def demo_fn(speech_upl: str, noise_type: str, snr: int, mic_input: str):
101
+ if mic_input:
102
+ speech_upl = mic_input
103
+ sr = config("sr", 48000, int, section="df")
104
+ logger.info(f"Got parameters speech_upl: {speech_upl}, noise: {noise_type}, snr: {snr}")
105
+ snr = int(snr)
106
+ noise_fn = NOISES[noise_type]
107
+ meta = AudioMetaData(-1, -1, -1, -1, "")
108
+ max_s = 10 # limit to 10 seconds
109
+ if speech_upl is not None:
110
+ sample, meta = load_audio(speech_upl, sr)
111
+ max_len = max_s * sr
112
+ if sample.shape[-1] > max_len:
113
+ start = torch.randint(0, sample.shape[-1] - max_len, ()).item()
114
+ sample = sample[..., start : start + max_len]
115
+ else:
116
+ sample, meta = load_audio("samples/p232_013_clean.wav", sr)
117
+ sample = sample[..., : max_s * sr]
118
+ if sample.dim() > 1 and sample.shape[0] > 1:
119
+ assert (
120
+ sample.shape[1] > sample.shape[0]
121
+ ), f"Expecting channels first, but got {sample.shape}"
122
+ sample = sample.mean(dim=0, keepdim=True)
123
+ logger.info(f"Loaded sample with shape {sample.shape}")
124
+ if noise_fn is not None:
125
+ noise, _ = load_audio(noise_fn, sr) # type: ignore
126
+ logger.info(f"Loaded noise with shape {noise.shape}")
127
+ _, _, sample = mix_at_snr(sample, noise, snr)
128
+ logger.info("Start denoising audio")
129
+ enhanced = enhance(model, df, sample)
130
+ logger.info("Denoising finished")
131
+ lim = torch.linspace(0.0, 1.0, int(sr * 0.15)).unsqueeze(0)
132
+ lim = torch.cat((lim, torch.ones(1, enhanced.shape[1] - lim.shape[1])), dim=1)
133
+ enhanced = enhanced * lim
134
+ if meta.sample_rate != sr:
135
+ enhanced = resample(enhanced, sr, meta.sample_rate)
136
+ sample = resample(sample, sr, meta.sample_rate)
137
+ sr = meta.sample_rate
138
+ noisy_wav = tempfile.NamedTemporaryFile(suffix="noisy.wav", delete=False).name
139
+ save_audio(noisy_wav, sample, sr)
140
+ enhanced_wav = tempfile.NamedTemporaryFile(suffix="enhanced.wav", delete=False).name
141
+ save_audio(enhanced_wav, enhanced, sr)
142
+ logger.info(f"saved audios: {noisy_wav}, {enhanced_wav}")
143
+ ax_noisy.clear()
144
+ ax_enh.clear()
145
+ noisy_im = spec_im(sample, sr=sr, figure=fig_noisy, ax=ax_noisy)
146
+ enh_im = spec_im(enhanced, sr=sr, figure=fig_enh, ax=ax_enh)
147
+ # noisy_wav = gr.make_waveform(noisy_fn, bar_count=200)
148
+ # enh_wav = gr.make_waveform(enhanced_fn, bar_count=200)
149
+ return noisy_wav, noisy_im, enhanced_wav, enh_im
150
+
151
+
152
+ def specshow(
153
+ spec,
154
+ ax=None,
155
+ title=None,
156
+ xlabel=None,
157
+ ylabel=None,
158
+ sr=48000,
159
+ n_fft=None,
160
+ hop=None,
161
+ t=None,
162
+ f=None,
163
+ vmin=-100,
164
+ vmax=0,
165
+ xlim=None,
166
+ ylim=None,
167
+ cmap="inferno",
168
+ ):
169
+ """Plots a spectrogram of shape [F, T]"""
170
+ spec_np = spec.cpu().numpy() if isinstance(spec, torch.Tensor) else spec
171
+ if ax is not None:
172
+ set_title = ax.set_title
173
+ set_xlabel = ax.set_xlabel
174
+ set_ylabel = ax.set_ylabel
175
+ set_xlim = ax.set_xlim
176
+ set_ylim = ax.set_ylim
177
+ else:
178
+ ax = plt
179
+ set_title = plt.title
180
+ set_xlabel = plt.xlabel
181
+ set_ylabel = plt.ylabel
182
+ set_xlim = plt.xlim
183
+ set_ylim = plt.ylim
184
+ if n_fft is None:
185
+ if spec.shape[0] % 2 == 0:
186
+ n_fft = spec.shape[0] * 2
187
+ else:
188
+ n_fft = (spec.shape[0] - 1) * 2
189
+ hop = hop or n_fft // 4
190
+ if t is None:
191
+ t = np.arange(0, spec_np.shape[-1]) * hop / sr
192
+ if f is None:
193
+ f = np.arange(0, spec_np.shape[0]) * sr // 2 / (n_fft // 2) / 1000
194
+ im = ax.pcolormesh(
195
+ t, f, spec_np, rasterized=True, shading="auto", vmin=vmin, vmax=vmax, cmap=cmap
196
+ )
197
+ if title is not None:
198
+ set_title(title)
199
+ if xlabel is not None:
200
+ set_xlabel(xlabel)
201
+ if ylabel is not None:
202
+ set_ylabel(ylabel)
203
+ if xlim is not None:
204
+ set_xlim(xlim)
205
+ if ylim is not None:
206
+ set_ylim(ylim)
207
+ return im
208
+
209
+
210
+ def spec_im(
211
+ audio: torch.Tensor,
212
+ figsize=(15, 5),
213
+ colorbar=False,
214
+ colorbar_format=None,
215
+ figure=None,
216
+ labels=True,
217
+ **kwargs,
218
+ ) -> Image:
219
+ audio = torch.as_tensor(audio)
220
+ if labels:
221
+ kwargs.setdefault("xlabel", "Time [s]")
222
+ kwargs.setdefault("ylabel", "Frequency [Hz]")
223
+ n_fft = kwargs.setdefault("n_fft", 1024)
224
+ hop = kwargs.setdefault("hop", 512)
225
+ w = torch.hann_window(n_fft, device=audio.device)
226
+ spec = torch.stft(audio, n_fft, hop, window=w, return_complex=False)
227
+ spec = spec.div_(w.pow(2).sum())
228
+ spec = torch.view_as_complex(spec).abs().clamp_min(1e-12).log10().mul(10)
229
+ kwargs.setdefault("vmax", max(0.0, spec.max().item()))
230
+
231
+ if figure is None:
232
+ figure = plt.figure(figsize=figsize)
233
+ figure.set_tight_layout(True)
234
+ if spec.dim() > 2:
235
+ spec = spec.squeeze(0)
236
+ im = specshow(spec, **kwargs)
237
+ if colorbar:
238
+ ckwargs = {}
239
+ if "ax" in kwargs:
240
+ if colorbar_format is None:
241
+ if kwargs.get("vmin", None) is not None or kwargs.get("vmax", None) is not None:
242
+ colorbar_format = "%+2.0f dB"
243
+ ckwargs = {"ax": kwargs["ax"]}
244
+ plt.colorbar(im, format=colorbar_format, **ckwargs)
245
+ figure.canvas.draw()
246
+ return Image.frombytes("RGB", figure.canvas.get_width_height(), figure.canvas.tostring_rgb())
247
+
248
+
249
+ def toggle(choice):
250
+ if choice == "mic":
251
+ return gr.update(visible=True, value=None), gr.update(visible=False, value=None)
252
+ else:
253
+ return gr.update(visible=False, value=None), gr.update(visible=True, value=None)
254
+
255
+
256
+ with gr.Blocks() as demo:
257
+ with gr.Row():
258
+ gr.Markdown(
259
+ """
260
+ ## DeepFilterNet2 Demo\
261
+
262
+ This demo denoises audio files using DeepFilterNet. Try it with your own voice!
263
+ """
264
+ )
265
+ with gr.Row():
266
+ with gr.Column():
267
+ radio = gr.Radio(
268
+ ["mic", "file"], value="file", label="How would you like to upload your audio?"
269
+ )
270
+ mic_input = gr.Mic(label="Input", type="filepath", visible=False)
271
+ audio_file = gr.Audio(type="filepath", label="Input", visible=True)
272
+ inputs = [
273
+ audio_file,
274
+ gr.Dropdown(
275
+ label="Add background noise",
276
+ choices=list(NOISES.keys()),
277
+ value="None",
278
+ ),
279
+ gr.Dropdown(
280
+ label="Noise Level (SNR)",
281
+ choices=["-5", "0", "10", "20"],
282
+ value="10",
283
+ ),
284
+ mic_input,
285
+ ]
286
+ btn = gr.Button("Generate")
287
+ with gr.Column():
288
+ outputs = [
289
+ # gr.Video(type="filepath", label="Noisy audio"),
290
+ gr.Audio(type="filepath", label="Noisy audio"),
291
+ gr.Image(label="Noisy spectrogram"),
292
+ # gr.Video(type="filepath", label="Enhanced audio"),
293
+ gr.Audio(type="filepath", label="Enhanced audio"),
294
+ gr.Image(label="Enhanced spectrogram"),
295
+ ]
296
+ btn.click(fn=demo_fn, inputs=inputs, outputs=outputs)
297
+ radio.change(toggle, radio, [mic_input, audio_file])
298
+ gr.Examples(
299
+ [
300
+ ["./samples/p232_013_clean.wav", "Kitchen", "10"],
301
+ ["./samples/p232_013_clean.wav", "Cafe", "10"],
302
+ ["./samples/p232_019_clean.wav", "Cafe", "10"],
303
+ ["./samples/p232_019_clean.wav", "River", "10"],
304
+ ],
305
+ fn=demo_fn,
306
+ inputs=inputs,
307
+ outputs=outputs,
308
+ cache_examples=True,
309
+ ),
310
+ gr.Markdown(open("usage.md").read())
311
+
312
+
313
+ demo.launch(enable_queue=True)
packages.txt ADDED
@@ -0,0 +1 @@
 
 
1
+ ffmpeg
pyproject.toml ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ [tool.black]
2
+ line-length = 100
3
+ target-version = ["py37", "py38", "py39", "py310"]
4
+ include = '\.pyi?$'
5
+
6
+ [tool.isort]
7
+ profile = "black"
8
+ line_length = 100
9
+ skip_gitignore = true
10
+ known_first_party = ["df", "libdf", "libdfdata"]
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ torch==1.13
2
+ torchaudio==0.13
3
+ deepfilternet==0.4.0
4
+ matplotlib==3.6
5
+ gradio==3.17
6
+ Pillow==9.3
samples/dkitchen.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:6bc229639249ce876bd40bbff15eae4553b7c15cdcf0b720ed814062b4956af6
3
+ size 2880044
samples/dliving.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:56156a12836c602a44b658a4a3f586e87909dd7b21402333fc777e0ad3c277e6
3
+ size 960044
samples/nriver.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:1aee2a2cae1f4f0d88f77c5d1616e6fbbc23bad88d287a21700e5b8519426e75
3
+ size 2880044
samples/p232_013_clean.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d7a51b4fdfb02657cf9410dbd34b4ea165acbec48581a8a074e1d45fdd3b3334
3
+ size 378612
samples/p232_019_clean.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2268caefff20ce9658a154c1481cc37ebf89347b3ef848b8a07115be3ec9c069
3
+ size 646658
samples/scafe.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:52cf963076f6de7d7c837ad5303b9b6416854558111ac5d6bd23926b38da281e
3
+ size 2880044
usage.md ADDED
@@ -0,0 +1,9 @@
 
 
 
 
 
 
 
 
 
 
1
+ **Usage:**
2
+
3
+ This demo takes a audio sample and enhances it using DeepFilterNet2.
4
+ Upload a (noisy) speech sample. You may optionally add some additional background noise to the input sample.
5
+ If no samples are provided, a default will be used.
6
+
7
+ Long audio files will be trimmed to 10s.
8
+
9
+ DeepFilterNet2 [(link)](https://github.com/Rikorose/DeepFilterNet) is used to denoise the noisy mixture.