File size: 5,363 Bytes
6dea0da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
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