BABE-2 / model /blind_bwe_utils.py
Vansh Chugh
initial deploy
6dea0da
Raw
History Blame Contribute Delete
5.36 kB
import torch
import numpy as np
def smooth_LTAS(LTAS, f, Noct=1):
# based on https://github.com/IoSR-Surrey/MatlabToolbox/blob/4bff1bb2da7c95de0ce2713e7c710a0afa70c705/%2Biosr/%2Bdsp/smoothSpectrum.m
def gauss_f(f_x, F, Noct):
sigma = (F / Noct) / np.pi
g = torch.exp(-(((f_x - F)**2) / (2 * (sigma**2))))
g = g / torch.sum(g)
return g
x_oct = LTAS.clone()
if Noct > 0:
for i in range(1, len(f)):
g = gauss_f(f, f[i], Noct)
g = g.to(LTAS.device)
x_oct[i] = torch.sum(g * LTAS)
if torch.all(LTAS >= 0):
x_oct[x_oct < 0] = 0
return x_oct
def compute_LTAS(x, sample_rate, nfft=2048, hop_length=512, win_length=2048, normalize=None, sqrt=False):
"""
Long-term average spectrum of a single waveform already at sample_rate.
"""
if len(x.shape) == 2:
x = np.mean(x, axis=1)
x = torch.tensor(x, dtype=torch.float32)
if normalize is not None:
std = x.std()
x = normalize * x / x.std()
else:
std = 1
X = torch.stft(x, n_fft=nfft, hop_length=hop_length, win_length=win_length, window=torch.hann_window(win_length), return_complex=True) / torch.sqrt(torch.hann_window(win_length).sum())
if sqrt:
Xsum = torch.sqrt(torch.sum(torch.abs(X)**2, dim=1).unsqueeze(-1))
else:
Xsum = torch.sum(torch.abs(X)**2, dim=1).unsqueeze(-1)
L = X.shape[-1]
X_norm = Xsum / L
X_mean = torch.mean(X_norm, dim=-1)
return X_mean, std
def apply_stft(x, NFFT):
window = torch.hamming_window(window_length=NFFT)
window = window.to(x.device)
x = torch.cat((x, torch.zeros(*x.shape[:-1], NFFT).to(x.device)), 1)
X = torch.stft(x, NFFT, hop_length=NFFT // 2, window=window, center=False, onesided=True, return_complex=True)
X = torch.view_as_real(X)
return X
def apply_filter_istft(X, H, NFFT):
window = torch.hamming_window(window_length=NFFT)
window = window.to(X.device)
X = X * H.unsqueeze(-1).unsqueeze(-1).expand(X.shape)
X = torch.view_as_complex(X)
x = torch.istft(X, NFFT, hop_length=NFFT // 2, window=window, center=False, return_complex=False)
return x
def apply_filter(x, H, NFFT):
X = apply_stft(x, NFFT)
xrec = apply_filter_istft(X, H, NFFT)
xrec = xrec[:, :x.shape[-1]]
return xrec
def design_filter_3(params, f, block_low_freq=False):
"""
Parametric shelving/EQ filter defined by a reference frequency (fref) and
piecewise log-linear slopes above (fc_p/A_p) and below (fc_m/A_m) it.
"""
fref = params[0]
fc_p = params[1]
fc_m = params[2]
A_p = params[3]
A_m = params[4]
assert (fc_p <= fref).any() == False, f"fc_p must be greater than fref: {fc_p}, {fref}"
assert (fc_m >= fref).any() == False, f"fc_m must be smaller than fre: {fc_m}, {fref}"
assert (fc_m <= f[1]).any() == False, f"fc_m must be greater than the minimum frequency: {fc_m}, {f[1]}"
assert (fc_p >= f[-1]).any() == False, f"fc_p must be smaller than the maximum frequency: {fc_p}, {f[-1]}"
f = f[1:]
H = torch.ones(f.shape).to(f.device)
H[f >= fref] = 10**(A_p[0] * torch.log2(f[f >= fref] / fref) / 20)
for i in range(0, len(fc_p)):
H[f >= fc_p[i]] = 10**(A_p[i + 1] * torch.log2(f[f >= fc_p[i]] / fc_p[i]) / 20) * H[f >= fc_p[i]][0]
if not block_low_freq:
H[f < fref] = 10**(A_m[0] * torch.log2(f[f < fref] / fref) / 20) * H[f < fref][-1]
for i in range(0, len(fc_m)):
H[f < fc_m[i]] = 10**(A_m[i + 1] * torch.log2(f[f < fc_m[i]] / fc_m[i]) / 20) * H[f < fc_m[i]][-1]
H = torch.cat((torch.zeros(1).to(f.device), H), 0)
return H
def apply_filter_and_norm_STFTmag_fweighted(X, Xref, H, freq_weight="linear"):
# X: (N,513, T) denoised estimate STFT, Xref: (N,513, T) observation STFT, H: (513,) filter
X = torch.sqrt(X[..., 0]**2 + X[..., 1]**2)
Xref = torch.sqrt(Xref[..., 0]**2 + Xref[..., 1]**2)
X = X * H.unsqueeze(-1).expand(X.shape)
freqs = torch.linspace(0, 1, X.shape[1]).to(X.device)
if freq_weight == "linear":
X = X * freqs.unsqueeze(-1).expand(X.shape)
Xref = Xref * freqs.unsqueeze(-1).expand(Xref.shape)
elif freq_weight == "None":
pass
elif freq_weight == "log":
X = X * torch.log2(1 + freqs.unsqueeze(-1).expand(X.shape))
Xref = Xref * torch.log2(1 + freqs.unsqueeze(-1).expand(Xref.shape))
elif freq_weight == "sqrt":
X = X * torch.sqrt(freqs.unsqueeze(-1).expand(X.shape))
Xref = Xref * torch.sqrt(freqs.unsqueeze(-1).expand(Xref.shape))
elif freq_weight == "log2":
X = X * torch.log2(freqs.unsqueeze(-1).expand(X.shape))
Xref = Xref * torch.log2(freqs.unsqueeze(-1).expand(Xref.shape))
elif freq_weight == "log10":
X = X * torch.log10(freqs.unsqueeze(-1).expand(X.shape))
Xref = Xref * torch.log10(freqs.unsqueeze(-1).expand(Xref.shape))
elif freq_weight == "cubic":
X = X * freqs.unsqueeze(-1).expand(X.shape)**3
Xref = Xref * freqs.unsqueeze(-1).expand(Xref.shape)**3
elif freq_weight == "quadratic":
X = X * freqs.unsqueeze(-1).expand(X.shape)**2
Xref = Xref * freqs.unsqueeze(-1).expand(Xref.shape)**2
norm = torch.linalg.norm(X.reshape(-1) - Xref.reshape(-1), ord=2)
return norm