smad / feature_extraction_smad.py
duclvQ's picture
Add Transformers load-by-id model files
64ef8a3 verified
Raw
History Blame Contribute Delete
2.36 kB
import numpy as np
from transformers import BatchFeature, SequenceFeatureExtractor
class SmadFeatureExtractor(SequenceFeatureExtractor):
model_input_names = ["input_features"]
def __init__(
self,
feature_size=80,
sampling_rate=16000,
padding_value=0.0,
segment_seconds=4.0,
n_fft=400,
hop_length=160,
**kwargs,
):
super().__init__(
feature_size=feature_size,
sampling_rate=sampling_rate,
padding_value=padding_value,
**kwargs,
)
self.segment_seconds = segment_seconds
self.n_fft = n_fft
self.hop_length = hop_length
def waveform_to_mel(self, waveform, sampling_rate=None):
import librosa
sampling_rate = sampling_rate or self.sampling_rate
mel = librosa.feature.melspectrogram(
y=np.asarray(waveform, dtype=np.float32),
sr=sampling_rate,
n_fft=self.n_fft,
hop_length=self.hop_length,
n_mels=self.feature_size,
power=2.0,
)
return librosa.power_to_db(mel).T.astype(np.float32)
def __call__(self, raw_speech, sampling_rate=None, return_tensors=None, **kwargs):
sampling_rate = sampling_rate or self.sampling_rate
if sampling_rate != self.sampling_rate:
raise ValueError(
f"Expected {self.sampling_rate} Hz audio. Resample before calling "
f"the feature extractor; received {sampling_rate} Hz."
)
if isinstance(raw_speech, np.ndarray) and raw_speech.ndim == 1:
waves = [raw_speech]
else:
waves = [np.asarray(w, dtype=np.float32) for w in raw_speech]
target_len = int(round(self.segment_seconds * self.sampling_rate))
features = []
for wave in waves:
if wave.ndim != 1:
raise ValueError("Expected mono audio arrays with shape `(samples,)`.")
if wave.shape[0] < target_len:
wave = np.pad(wave, (0, target_len - wave.shape[0]))
elif wave.shape[0] > target_len:
wave = wave[:target_len]
features.append(self.waveform_to_mel(wave, sampling_rate=sampling_rate))
return BatchFeature({"input_features": features}, tensor_type=return_tensors)