File size: 2,848 Bytes
320e2b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import json
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchaudio
import soundfile as sf

TARGET_SR = 44100
TARGET_DURATION = 5
TARGET_NUM_SAMPLES = TARGET_SR * TARGET_DURATION


class NormalizeMeanStd(nn.Module):
    def __init__(self, mean: float, std: float, eps: float = 1e-6):
        super().__init__()
        self.register_buffer("mean", torch.tensor(mean).view(1, 1, 1))
        self.register_buffer("std", torch.tensor(std).view(1, 1, 1))
        self.eps = eps

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        return (x - self.mean) / (self.std + self.eps)


def mel_transform_from_stats(

    stats_path: str,

    sample_rate: int = 44100,

    n_fft: int = 1024,

    hop_length: int = 256,

    n_mels: int = 40,

):
    with open(stats_path, "r", encoding="utf-8") as f:
        s = json.load(f)

    return nn.Sequential(
        torchaudio.transforms.MelSpectrogram(
            sample_rate=sample_rate,
            n_fft=n_fft,
            hop_length=hop_length,
            n_mels=n_mels,
        ),
        torchaudio.transforms.AmplitudeToDB(),
        NormalizeMeanStd(s["mean"], s["std"]),
    )


def load_audio_with_soundfile(file_path: str):
    waveform, sr = sf.read(file_path, always_2d=True)   # [samples, channels]
    waveform = torch.tensor(waveform, dtype=torch.float32).transpose(0, 1)  # [channels, samples]
    return waveform, sr


def convert_to_mono(waveform: torch.Tensor) -> torch.Tensor:
    if waveform.shape[0] > 1:
        waveform = waveform.mean(dim=0, keepdim=True)
    return waveform


def resample_if_needed(waveform: torch.Tensor, orig_sr: int, target_sr: int = TARGET_SR) -> torch.Tensor:
    if orig_sr != target_sr:
        resampler = torchaudio.transforms.Resample(orig_freq=orig_sr, new_freq=target_sr)
        waveform = resampler(waveform)
    return waveform


def pad_or_trim(waveform: torch.Tensor, target_num_samples: int = TARGET_NUM_SAMPLES) -> torch.Tensor:
    num_samples = waveform.shape[1]

    if num_samples > target_num_samples:
        waveform = waveform[:, :target_num_samples]
    elif num_samples < target_num_samples:
        pad_amount = target_num_samples - num_samples
        waveform = F.pad(waveform, (0, pad_amount))

    return waveform


def preprocess_audio(file_path: str, stats_path: str) -> torch.Tensor:
    waveform, sr = load_audio_with_soundfile(file_path)

    waveform = convert_to_mono(waveform)
    waveform = resample_if_needed(waveform, sr, TARGET_SR)
    waveform = pad_or_trim(waveform, TARGET_NUM_SAMPLES)

    mel_transform = mel_transform_from_stats(stats_path=stats_path)
    features = mel_transform(waveform)   # [1, 40, time]

    features = features.unsqueeze(0)     # [1, 1, 40, time]
    return features