File size: 4,402 Bytes
7c268e9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""NeMo-compatible log-mel front-end (numpy only, no torch/librosa at runtime).

Mirrors ``FilterbankFeatures`` (eval mode) from NeMo Speech 3.0 with the
Sortformer streaming configuration:

    sample_rate=16000, n_fft=512, win_length=400 (hann, periodic=False),
    hop_length=160, n_mels=128, preemph=0.97, mag_power=2.0,
    log(x + 2**-24), normalize="NA", pad_to=16

``log_mel_spectrogram`` returns ``(num_frames, 128)`` float32 where
``num_frames = floor(num_samples / hop)`` rounded up to a multiple of 16.
"""

import numpy as np

SAMPLE_RATE = 16000
N_FFT = 512
WIN_LENGTH = 400
HOP_LENGTH = 160
N_MELS = 128
PREEMPH = 0.97
LOG_ZERO_GUARD = 2.0 ** -24
PAD_TO = 16

_MEL_CACHE = {}


def _hz_to_mel(freq):
    # Slaney mel scale (librosa default, htk=False)
    f_sp = 200.0 / 3.0
    log_step = np.log(6.4) / 27.0
    min_log_hz = 1000.0
    min_log_mel = min_log_hz / f_sp
    freq = np.asarray(freq, dtype=np.float64)
    mel_linear = freq / f_sp
    mel_log = min_log_mel + np.log(np.maximum(freq, min_log_hz) / min_log_hz) / log_step
    return np.where(freq >= min_log_hz, mel_log, mel_linear)


def _mel_to_hz(mel):
    f_sp = 200.0 / 3.0
    log_step = np.log(6.4) / 27.0
    min_log_hz = 1000.0
    min_log_mel = min_log_hz / f_sp
    mel = np.asarray(mel, dtype=np.float64)
    hz_linear = f_sp * mel
    hz_log = min_log_hz * np.exp(log_step * (mel - min_log_mel))
    return np.where(mel >= min_log_mel, hz_log, hz_linear)


def mel_filterbank(sample_rate: int = SAMPLE_RATE, n_fft: int = N_FFT, n_mels: int = N_MELS):
    """librosa-compatible Slaney-normalized Slaney-scale mel filterbank."""
    key = (sample_rate, n_fft, n_mels)
    if key in _MEL_CACHE:
        return _MEL_CACHE[key]

    fmin, fmax = 0.0, sample_rate / 2.0
    mels = np.linspace(_hz_to_mel(fmin), _hz_to_mel(fmax), n_mels + 2)
    hz = _mel_to_hz(mels)
    freqs = np.linspace(0.0, sample_rate / 2.0, 1 + n_fft // 2)

    fdiff = np.diff(hz)
    ramps = np.subtract.outer(hz, freqs)
    lower = -ramps[np.arange(n_mels), :] / fdiff[np.arange(n_mels)][:, None]
    upper = ramps[np.arange(2, n_mels + 2), :] / fdiff[np.arange(1, n_mels + 1)][:, None]
    weights = np.maximum(0.0, np.minimum(lower, upper))

    enorm = 2.0 / (hz[2 : n_mels + 2] - hz[:n_mels])
    weights *= enorm[:, None]
    result = weights.astype(np.float64)
    _MEL_CACHE[key] = result
    return result


def _hann_window(length: int):
    return 0.5 - 0.5 * np.cos(2.0 * np.pi * np.arange(length) / (length - 1))


def log_mel_spectrogram(
    waveform: np.ndarray,
    sample_rate: int = SAMPLE_RATE,
    *,
    preemph: float = PREEMPH,
    pad_to: int = PAD_TO,
    mel_filter: np.ndarray = None,
) -> np.ndarray:
    """Compute the NeMo Sortformer streaming log-mel features for a mono waveform."""
    if sample_rate != SAMPLE_RATE:
        raise ValueError(f"expected {SAMPLE_RATE} Hz input, got {sample_rate}")
    waveform = np.asarray(waveform, dtype=np.float32).reshape(-1)

    if preemph is not None and waveform.size > 0:
        emphasized = np.empty_like(waveform)
        emphasized[0] = waveform[0]
        emphasized[1:] = waveform[1:] - preemph * waveform[:-1]
        waveform = emphasized

    padded = np.pad(waveform, (N_FFT // 2, N_FFT // 2), mode="constant")
    num_samples = waveform.shape[0]
    num_frames = num_samples // HOP_LENGTH
    if num_frames == 0:
        return np.zeros((0, N_MELS), dtype=np.float32)

    window = _hann_window(WIN_LENGTH).astype(np.float32)
    left_pad = (N_FFT - WIN_LENGTH) // 2
    frame_starts = np.arange(num_frames) * HOP_LENGTH
    frames = np.lib.stride_tricks.as_strided(
        padded,
        shape=(num_frames, N_FFT),
        strides=(padded.strides[0] * HOP_LENGTH, padded.strides[0]),
        writeable=False,
    ).copy()
    frames[:, :left_pad] = 0.0
    frames[:, left_pad : left_pad + WIN_LENGTH] *= window
    frames[:, left_pad + WIN_LENGTH :] = 0.0

    spectrum = np.fft.rfft(frames, n=N_FFT, axis=1)
    magnitude = np.abs(spectrum).astype(np.float32) ** 2.0

    if mel_filter is None:
        mel_filter = mel_filterbank()
    mel = magnitude @ mel_filter.T.astype(np.float32)
    log_mel = np.log(mel + LOG_ZERO_GUARD, dtype=np.float32)

    if pad_to and log_mel.shape[0] % pad_to:
        pad = pad_to - log_mel.shape[0] % pad_to
        log_mel = np.pad(log_mel, ((0, pad), (0, 0)), mode="constant")
    return log_mel.astype(np.float32)