authtransforms 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- authtransforms/__init__.py +65 -0
- authtransforms/pipeline.py +196 -0
- authtransforms/transforms.py +648 -0
- authtransforms/utils.py +364 -0
- authtransforms-0.1.0.dist-info/METADATA +438 -0
- authtransforms-0.1.0.dist-info/RECORD +8 -0
- authtransforms-0.1.0.dist-info/WHEEL +4 -0
- authtransforms-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,648 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Audio augmentation transforms for PyTorch, inspired by:
|
|
3
|
+
https://jonathanbgn.com/2021/08/30/audio-augmentation.html
|
|
4
|
+
|
|
5
|
+
Transforms follow the torchvision convention: callable objects with __call__
|
|
6
|
+
that accept a (channels, samples) tensor and return a transformed tensor.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import math
|
|
10
|
+
import os
|
|
11
|
+
import pathlib
|
|
12
|
+
import random
|
|
13
|
+
|
|
14
|
+
import torch
|
|
15
|
+
import torchaudio
|
|
16
|
+
import torchaudio.transforms as T
|
|
17
|
+
import tempfile
|
|
18
|
+
import random
|
|
19
|
+
import subprocess
|
|
20
|
+
import scipy.signal
|
|
21
|
+
import numpy as np
|
|
22
|
+
from pathlib import Path
|
|
23
|
+
from scipy.signal import fftconvolve
|
|
24
|
+
import numpy as np
|
|
25
|
+
import librosa
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
# ---------------------------------------------------------------------------
|
|
30
|
+
# Core transforms
|
|
31
|
+
# ---------------------------------------------------------------------------
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class RandomClip:
|
|
35
|
+
"""Extract a random fixed-length clip from the audio.
|
|
36
|
+
|
|
37
|
+
Args:
|
|
38
|
+
sample_rate: Sample rate of the audio.
|
|
39
|
+
clip_length: Desired clip length in samples.
|
|
40
|
+
vad: If True, apply Voice Activity Detection to trim leading/trailing
|
|
41
|
+
silence after clipping.
|
|
42
|
+
vad_trigger_level: VAD sensitivity (higher = more aggressive trimming).
|
|
43
|
+
"""
|
|
44
|
+
|
|
45
|
+
def __init__(self, sample_rate: int, clip_length: int, vad: bool = True, vad_trigger_level: float = 7.0):
|
|
46
|
+
self.clip_length = clip_length
|
|
47
|
+
self.vad = T.Vad(sample_rate=sample_rate, trigger_level=vad_trigger_level) if vad else None
|
|
48
|
+
#VAD is Voice Activity Detection
|
|
49
|
+
|
|
50
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
51
|
+
audio_length = audio.shape[-1]
|
|
52
|
+
if audio_length > self.clip_length:
|
|
53
|
+
offset = random.randint(0, audio_length - self.clip_length) #to select random int position where offset is the int
|
|
54
|
+
audio = audio[..., offset : offset + self.clip_length] #Torch slicing
|
|
55
|
+
if self.vad is not None:
|
|
56
|
+
audio = self.vad(audio)
|
|
57
|
+
return audio
|
|
58
|
+
|
|
59
|
+
def __repr__(self) -> str:
|
|
60
|
+
return f"{self.__class__.__name__}(clip_length={self.clip_length})"
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
class RandomSpeedChange:
|
|
64
|
+
"""Randomly perturb audio playback speed.
|
|
65
|
+
|
|
66
|
+
Resamples back to the original sample rate so the output length stays
|
|
67
|
+
approximately the same. Uses torchaudio's SoX effects backend.
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
sample_rate: Sample rate of the audio.
|
|
71
|
+
speed_factors: Sequence of speed multipliers to choose from.
|
|
72
|
+
Default is (0.9, 1.0, 1.1) as used in the literature.
|
|
73
|
+
"""
|
|
74
|
+
|
|
75
|
+
def __init__(self, sample_rate: int, speed_factors=(0.9, 1.0, 1.1)):
|
|
76
|
+
self.sample_rate = sample_rate
|
|
77
|
+
self.speed_factors = list(speed_factors)
|
|
78
|
+
|
|
79
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
80
|
+
speed_factor = random.choice(self.speed_factors)
|
|
81
|
+
if speed_factor == 1.0:
|
|
82
|
+
return audio
|
|
83
|
+
|
|
84
|
+
sox_effects = [ #This is the sox effect that changes the speed and sample_rate of the audio
|
|
85
|
+
["speed", str(speed_factor)],
|
|
86
|
+
["rate", str(self.sample_rate)],
|
|
87
|
+
]
|
|
88
|
+
transformed, _ = torchaudio.sox_effects.apply_effects_tensor(
|
|
89
|
+
audio, self.sample_rate, sox_effects
|
|
90
|
+
)
|
|
91
|
+
return transformed
|
|
92
|
+
|
|
93
|
+
def __repr__(self) -> str:
|
|
94
|
+
return f"{self.__class__.__name__}(speed_factors={self.speed_factors})"
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class RandomBackgroundNoise:
|
|
98
|
+
"""Mix the audio with a random noise file at a random SNR level.
|
|
99
|
+
|
|
100
|
+
Args:
|
|
101
|
+
sample_rate: Target sample rate.
|
|
102
|
+
noise_dir: Directory (searched recursively) containing .wav noise files.
|
|
103
|
+
min_snr_db: Minimum signal-to-noise ratio in dB.
|
|
104
|
+
max_snr_db: Maximum signal-to-noise ratio in dB.
|
|
105
|
+
"""
|
|
106
|
+
|
|
107
|
+
def __init__(self, sample_rate: int, noise_dir: str, min_snr_db: float = 0, max_snr_db: float = 15):
|
|
108
|
+
self.sample_rate = sample_rate
|
|
109
|
+
self.min_snr_db = min_snr_db
|
|
110
|
+
self.max_snr_db = max_snr_db
|
|
111
|
+
|
|
112
|
+
noise_path = pathlib.Path(noise_dir)
|
|
113
|
+
if not noise_path.exists(): #To check if the noise path exists
|
|
114
|
+
raise IOError(f"Noise directory `{noise_dir}` does not exist")
|
|
115
|
+
self.noise_files = list(noise_path.glob("**/*.wav")) # Find all the files inside the noise path
|
|
116
|
+
if not self.noise_files:
|
|
117
|
+
raise IOError(f"No .wav files found in `{noise_dir}`")
|
|
118
|
+
|
|
119
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
120
|
+
noise_file = random.choice(self.noise_files)
|
|
121
|
+
effects = [
|
|
122
|
+
["remix", "1"], # convert to mono (noise)
|
|
123
|
+
["rate", str(self.sample_rate)], # resample (noise)
|
|
124
|
+
]
|
|
125
|
+
noise, _ = torchaudio.sox_effects.apply_effects_file(noise_file, effects, normalize=True)
|
|
126
|
+
|
|
127
|
+
audio_length = audio.shape[-1]
|
|
128
|
+
noise_length = noise.shape[-1]
|
|
129
|
+
|
|
130
|
+
if noise_length > audio_length:
|
|
131
|
+
offset = random.randint(0, noise_length - audio_length) #To make the noise length equal to audio length
|
|
132
|
+
noise = noise[..., offset : offset + audio_length]
|
|
133
|
+
elif noise_length < audio_length:
|
|
134
|
+
# tile noise to cover the audio
|
|
135
|
+
repeats = math.ceil(audio_length / noise_length)
|
|
136
|
+
noise = noise.repeat(1, repeats)[..., :audio_length]
|
|
137
|
+
|
|
138
|
+
snr_db = random.uniform(self.min_snr_db, self.max_snr_db) #The snr defines whether the noise or audio is gonna dominate , high snr signal dominates, whereas low snr noise dominates
|
|
139
|
+
snr = math.exp(snr_db / 10)
|
|
140
|
+
audio_power = audio.norm(p=2) #Computing both the audio and noise power
|
|
141
|
+
noise_power = noise.norm(p=2)
|
|
142
|
+
scale = snr * noise_power / (audio_power + 1e-9)
|
|
143
|
+
|
|
144
|
+
return (scale * audio + noise) / 2 #Mix the audio and noise
|
|
145
|
+
|
|
146
|
+
def __repr__(self) -> str:
|
|
147
|
+
return (
|
|
148
|
+
f"{self.__class__.__name__}("
|
|
149
|
+
f"snr=[{self.min_snr_db}, {self.max_snr_db}] dB, "
|
|
150
|
+
f"n_files={len(self.noise_files)})"
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
class RandomPitchShift:
|
|
155
|
+
|
|
156
|
+
# This changes the depth of the voice , variation in male/female/child voices
|
|
157
|
+
"""Shift audio pitch by a random number of semitones.
|
|
158
|
+
|
|
159
|
+
Args:
|
|
160
|
+
sample_rate: Sample rate of the audio.
|
|
161
|
+
semitones: Range tuple (min, max) of semitone shifts.
|
|
162
|
+
"""
|
|
163
|
+
|
|
164
|
+
def __init__(self, sample_rate: int, semitones: tuple = (-2, 2)): #Semitone is the range of pitchshift
|
|
165
|
+
self.sample_rate = sample_rate
|
|
166
|
+
self.semitones = semitones
|
|
167
|
+
|
|
168
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
169
|
+
n = random.uniform(*self.semitones)
|
|
170
|
+
if n == 0:
|
|
171
|
+
return audio
|
|
172
|
+
sox_effects = [["pitch", str(int(n * 100))], ["rate", str(self.sample_rate)]] #1 semitone = 100 cents and the sample rate is resampled to original
|
|
173
|
+
transformed, _ = torchaudio.sox_effects.apply_effects_tensor(
|
|
174
|
+
audio, self.sample_rate, sox_effects
|
|
175
|
+
)
|
|
176
|
+
return transformed
|
|
177
|
+
|
|
178
|
+
def __repr__(self) -> str:
|
|
179
|
+
return f"{self.__class__.__name__}(semitones={self.semitones})"
|
|
180
|
+
|
|
181
|
+
|
|
182
|
+
class RandomGain:
|
|
183
|
+
"""Scale audio amplitude by a random gain factor (in dB).
|
|
184
|
+
|
|
185
|
+
Args:
|
|
186
|
+
min_gain_db: Minimum gain in dB.
|
|
187
|
+
max_gain_db: Maximum gain in dB.
|
|
188
|
+
"""
|
|
189
|
+
|
|
190
|
+
def __init__(self, min_gain_db: float = -6, max_gain_db: float = 6):
|
|
191
|
+
self.min_gain_db = min_gain_db
|
|
192
|
+
self.max_gain_db = max_gain_db
|
|
193
|
+
|
|
194
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
195
|
+
gain_db = random.uniform(self.min_gain_db, self.max_gain_db) # This is to make the audio louder or quieter
|
|
196
|
+
gain = 10 ** (gain_db / 20 ) #This is the formula that changes the amplitude to linear scales and multiply with audio.
|
|
197
|
+
return audio * gain
|
|
198
|
+
|
|
199
|
+
def __repr__(self) -> str:
|
|
200
|
+
return f"{self.__class__.__name__}(gain_db=[{self.min_gain_db}, {self.max_gain_db}])"
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
class AddGaussianNoise:
|
|
204
|
+
"""Add Gaussian white noise at a random amplitude.
|
|
205
|
+
|
|
206
|
+
Args:
|
|
207
|
+
min_amplitude: Minimum noise standard deviation.
|
|
208
|
+
max_amplitude: Maximum noise standard deviation.
|
|
209
|
+
"""
|
|
210
|
+
|
|
211
|
+
def __init__(self, min_amplitude: float = 0.001, max_amplitude: float = 0.015):
|
|
212
|
+
self.min_amplitude = min_amplitude #min strength of noise
|
|
213
|
+
self.max_amplitude = max_amplitude #max strength of noise
|
|
214
|
+
|
|
215
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
216
|
+
amplitude = random.uniform(self.min_amplitude, self.max_amplitude)
|
|
217
|
+
noise = torch.randn_like(audio) * amplitude #randn_like generates random values
|
|
218
|
+
return audio + noise
|
|
219
|
+
|
|
220
|
+
def __repr__(self) -> str:
|
|
221
|
+
return (
|
|
222
|
+
f"{self.__class__.__name__}("
|
|
223
|
+
f"amplitude=[{self.min_amplitude}, {self.max_amplitude}])"
|
|
224
|
+
)
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
class TimeShift:
|
|
228
|
+
"""Shift the audio waveform in time, wrapping or zero-padding.
|
|
229
|
+
|
|
230
|
+
Args:
|
|
231
|
+
max_shift: Maximum shift as a fraction of the total length (0–1).
|
|
232
|
+
roll: If True, wrap the audio (circular shift). If False, zero-pad.
|
|
233
|
+
"""
|
|
234
|
+
|
|
235
|
+
def __init__(self, max_shift: float = 0.2, roll: bool = False):
|
|
236
|
+
self.max_shift = max_shift
|
|
237
|
+
self.roll = roll
|
|
238
|
+
|
|
239
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
240
|
+
length = audio.shape[-1]
|
|
241
|
+
shift = random.randint(-int(length * self.max_shift), int(length * self.max_shift))
|
|
242
|
+
if shift == 0:
|
|
243
|
+
return audio
|
|
244
|
+
if self.roll:
|
|
245
|
+
return torch.roll(audio, shift, dims=-1)
|
|
246
|
+
# zero-pad version
|
|
247
|
+
result = torch.zeros_like(audio)
|
|
248
|
+
if shift > 0:
|
|
249
|
+
result[..., shift:] = audio[..., : length - shift]
|
|
250
|
+
else:
|
|
251
|
+
result[..., : length + shift] = audio[..., -shift:]
|
|
252
|
+
return result
|
|
253
|
+
|
|
254
|
+
def __repr__(self) -> str:
|
|
255
|
+
return f"{self.__class__.__name__}(max_shift={self.max_shift}, roll={self.roll})"
|
|
256
|
+
|
|
257
|
+
|
|
258
|
+
class SpecAugment:
|
|
259
|
+
"""Apply SpecAugment on the spectrogram (time & frequency masking).
|
|
260
|
+
|
|
261
|
+
This transform expects a raw audio tensor and internally computes a
|
|
262
|
+
spectrogram, applies masking, then returns the masked spectrogram.
|
|
263
|
+
Use it as the final step in a pipeline when your model consumes spectrograms.
|
|
264
|
+
|
|
265
|
+
Args:
|
|
266
|
+
sample_rate: Sample rate.
|
|
267
|
+
n_fft: FFT size.
|
|
268
|
+
n_mels: Number of mel filterbanks.
|
|
269
|
+
freq_mask_param: Maximum frequency mask width.
|
|
270
|
+
time_mask_param: Maximum time mask width.
|
|
271
|
+
n_freq_masks: Number of frequency masks.
|
|
272
|
+
n_time_masks: Number of time masks.
|
|
273
|
+
"""
|
|
274
|
+
|
|
275
|
+
def __init__(
|
|
276
|
+
self,
|
|
277
|
+
sample_rate: int,
|
|
278
|
+
n_fft: int = 400,
|
|
279
|
+
n_mels: int = 80,
|
|
280
|
+
freq_mask_param: int = 27,
|
|
281
|
+
time_mask_param: int = 100,
|
|
282
|
+
n_freq_masks: int = 2,
|
|
283
|
+
n_time_masks: int = 2,
|
|
284
|
+
):
|
|
285
|
+
self.mel_spec = T.MelSpectrogram(sample_rate=sample_rate, n_fft=n_fft, n_mels=n_mels)
|
|
286
|
+
self.freq_masking = T.FrequencyMasking(freq_mask_param=freq_mask_param)
|
|
287
|
+
self.time_masking = T.TimeMasking(time_mask_param=time_mask_param)
|
|
288
|
+
self.n_freq_masks = n_freq_masks
|
|
289
|
+
self.n_time_masks = n_time_masks
|
|
290
|
+
|
|
291
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
292
|
+
spec = self.mel_spec(audio)
|
|
293
|
+
for _ in range(self.n_freq_masks):
|
|
294
|
+
spec = self.freq_masking(spec)
|
|
295
|
+
for _ in range(self.n_time_masks):
|
|
296
|
+
spec = self.time_masking(spec)
|
|
297
|
+
return spec
|
|
298
|
+
|
|
299
|
+
def __repr__(self) -> str:
|
|
300
|
+
return f"{self.__class__.__name__}(n_freq={self.n_freq_masks}, n_time={self.n_time_masks})"
|
|
301
|
+
|
|
302
|
+
|
|
303
|
+
class RandomApply:
|
|
304
|
+
"""Apply a transform with probability p (mirrors torchvision.transforms.RandomApply).
|
|
305
|
+
|
|
306
|
+
Args:
|
|
307
|
+
transform: A callable transform.
|
|
308
|
+
p: Probability of applying the transform (0–1).
|
|
309
|
+
"""
|
|
310
|
+
|
|
311
|
+
def __init__(self, transform, p: float = 0.5):
|
|
312
|
+
self.transform = transform
|
|
313
|
+
self.p = p
|
|
314
|
+
|
|
315
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
316
|
+
if random.random() < self.p:
|
|
317
|
+
return self.transform(audio)
|
|
318
|
+
return audio
|
|
319
|
+
|
|
320
|
+
def __repr__(self) -> str:
|
|
321
|
+
return f"{self.__class__.__name__}(p={self.p}, transform={self.transform})"
|
|
322
|
+
|
|
323
|
+
|
|
324
|
+
class Normalize:
|
|
325
|
+
"""Normalize audio to peak amplitude of 1.0 (or a target level).
|
|
326
|
+
|
|
327
|
+
Args:
|
|
328
|
+
target_db: Target peak level in dBFS. Default -3 dBFS. Use None for 0 dBFS.
|
|
329
|
+
"""
|
|
330
|
+
|
|
331
|
+
def __init__(self, target_db: float = -3.0):
|
|
332
|
+
self.target_level = 10 ** (target_db / 20)
|
|
333
|
+
|
|
334
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
335
|
+
peak = audio.abs().max()
|
|
336
|
+
if peak < 1e-9:
|
|
337
|
+
return audio
|
|
338
|
+
return audio / peak * self.target_level
|
|
339
|
+
|
|
340
|
+
def __repr__(self) -> str:
|
|
341
|
+
return f"{self.__class__.__name__}(target_db={20 * math.log10(self.target_level):.1f})"
|
|
342
|
+
|
|
343
|
+
|
|
344
|
+
class ToMono:
|
|
345
|
+
"""Convert multi-channel audio to mono by averaging channels."""
|
|
346
|
+
|
|
347
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
348
|
+
if audio.shape[0] > 1:
|
|
349
|
+
return audio.mean(dim=0, keepdim=True)
|
|
350
|
+
return audio
|
|
351
|
+
|
|
352
|
+
def __repr__(self) -> str:
|
|
353
|
+
return f"{self.__class__.__name__}()"
|
|
354
|
+
|
|
355
|
+
class CodecAugmentation:
|
|
356
|
+
"""Apply random lossy codec compression (AAC, MP3, Opus, Vorbis, G.722).
|
|
357
|
+
|
|
358
|
+
Args:
|
|
359
|
+
sample_rate: Sample rate of the audio.
|
|
360
|
+
codecs: Sequence of codecs to sample from.
|
|
361
|
+
p: Probability of applying augmentation.
|
|
362
|
+
bitrate_range: Sequence of bitrates (bps) to choose from.
|
|
363
|
+
mp3_vbr_prob: Probability of using VBR for MP3.
|
|
364
|
+
"""
|
|
365
|
+
|
|
366
|
+
def __init__(
|
|
367
|
+
self,
|
|
368
|
+
sample_rate: int,
|
|
369
|
+
codecs: tuple = ("aac", "mp3", "opus", "vorbis","g722"),
|
|
370
|
+
p: float = 0.5,
|
|
371
|
+
bitrate_range: tuple = (12000, 16000, 24000, 32000),
|
|
372
|
+
mp3_vbr_prob: float = 0.3,
|
|
373
|
+
):
|
|
374
|
+
self.sample_rate = sample_rate
|
|
375
|
+
self.codecs = list(codecs)
|
|
376
|
+
self.p = p
|
|
377
|
+
self.bitrate_range = list(bitrate_range)
|
|
378
|
+
self.mp3_vbr_prob = mp3_vbr_prob
|
|
379
|
+
|
|
380
|
+
self.codec_map = {
|
|
381
|
+
"aac": ("aac", ".m4a"),
|
|
382
|
+
"mp3": ("libmp3lame", ".mp3"),
|
|
383
|
+
"opus": ("libopus", ".ogg"),
|
|
384
|
+
"vorbis": ("libvorbis", ".ogg"),
|
|
385
|
+
"g722": ("g722", ".g722"),
|
|
386
|
+
}
|
|
387
|
+
|
|
388
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
389
|
+
if random.random() > self.p:
|
|
390
|
+
return audio
|
|
391
|
+
|
|
392
|
+
codec = random.choice(self.codecs)
|
|
393
|
+
encoder, ext = self.codec_map[codec]
|
|
394
|
+
bitrate = random.choice(self.bitrate_range)
|
|
395
|
+
target_sr = 8000 if codec == "g722" else self.sample_rate
|
|
396
|
+
|
|
397
|
+
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as src, \
|
|
398
|
+
tempfile.NamedTemporaryFile(suffix=ext, delete=False) as enc, \
|
|
399
|
+
tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as dec:
|
|
400
|
+
|
|
401
|
+
try:
|
|
402
|
+
torchaudio.save(src.name, audio, self.sample_rate)
|
|
403
|
+
|
|
404
|
+
# Encode with ffmpeg
|
|
405
|
+
cmd = [
|
|
406
|
+
"ffmpeg", "-y",
|
|
407
|
+
"-i", src.name,
|
|
408
|
+
"-ar", str(target_sr),
|
|
409
|
+
"-c:a", encoder,
|
|
410
|
+
]
|
|
411
|
+
|
|
412
|
+
if codec == "mp3" and random.random() < self.mp3_vbr_prob:
|
|
413
|
+
cmd += ["-q:a", str(random.randint(2, 6))]
|
|
414
|
+
else:
|
|
415
|
+
cmd += ["-b:a", str(bitrate)]
|
|
416
|
+
|
|
417
|
+
cmd.append(enc.name)
|
|
418
|
+
subprocess.run(cmd, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)
|
|
419
|
+
|
|
420
|
+
# Decode back to WAV
|
|
421
|
+
subprocess.run([
|
|
422
|
+
"ffmpeg", "-y",
|
|
423
|
+
"-i", enc.name,
|
|
424
|
+
"-ar", str(self.sample_rate),
|
|
425
|
+
dec.name
|
|
426
|
+
], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True)
|
|
427
|
+
|
|
428
|
+
augmented, _ = torchaudio.load(dec.name)
|
|
429
|
+
return augmented
|
|
430
|
+
|
|
431
|
+
finally:
|
|
432
|
+
for f in [src.name, enc.name, dec.name]:
|
|
433
|
+
if os.path.exists(f):
|
|
434
|
+
os.remove(f)
|
|
435
|
+
|
|
436
|
+
def __repr__(self) -> str:
|
|
437
|
+
return f"{self.__class__.__name__}(sample_rate={self.sample_rate}, codecs={self.codecs}, p={self.p})"
|
|
438
|
+
|
|
439
|
+
class FilterAugmentation:
|
|
440
|
+
"""Apply a band-pass filter to simulate speech frequency range.
|
|
441
|
+
Args:
|
|
442
|
+
sample_rate: Sample rate of the audio.
|
|
443
|
+
low_hz: Lower cutoff frequency in Hz.
|
|
444
|
+
high_hz: Upper cutoff frequency in Hz.
|
|
445
|
+
p: Probability of applying augmentation.
|
|
446
|
+
"""
|
|
447
|
+
def __init__(
|
|
448
|
+
self,
|
|
449
|
+
sample_rate: int,
|
|
450
|
+
low_hz: float = 50.0,
|
|
451
|
+
high_hz: float = 7000.0,
|
|
452
|
+
p: float = 0.5,
|
|
453
|
+
):
|
|
454
|
+
self.sample_rate = sample_rate
|
|
455
|
+
self.low_hz = low_hz
|
|
456
|
+
self.high_hz = high_hz
|
|
457
|
+
self.p = p
|
|
458
|
+
|
|
459
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
460
|
+
if random.random() > self.p:
|
|
461
|
+
return audio
|
|
462
|
+
|
|
463
|
+
# Normalise cutoffs to [0, 1] where 1 = Nyquist frequency
|
|
464
|
+
nyquist = self.sample_rate / 2
|
|
465
|
+
low = self.low_hz / nyquist
|
|
466
|
+
high = self.high_hz / nyquist
|
|
467
|
+
|
|
468
|
+
# Design a Butterworth band-pass filter
|
|
469
|
+
sos = scipy.signal.butter(
|
|
470
|
+
N=4, # filter order
|
|
471
|
+
Wn=[low, high],
|
|
472
|
+
btype="bandpass",
|
|
473
|
+
output="sos" # numerically stable form
|
|
474
|
+
)
|
|
475
|
+
|
|
476
|
+
# Apply filter to each channel
|
|
477
|
+
audio_np = audio.numpy()
|
|
478
|
+
filtered = scipy.signal.sosfiltfilt(sos, audio_np) # zero-phase filtering
|
|
479
|
+
filtered = filtered.copy()
|
|
480
|
+
return torch.from_numpy(filtered).float()
|
|
481
|
+
|
|
482
|
+
def __repr__(self) -> str:
|
|
483
|
+
return (f"{self.__class__.__name__}("
|
|
484
|
+
f"low_hz={self.low_hz}, high_hz={self.high_hz}, p={self.p})")
|
|
485
|
+
|
|
486
|
+
|
|
487
|
+
class RoomImpulseResponse:
|
|
488
|
+
"""
|
|
489
|
+
Convolve audio with a randomly sampled Room Impulse Response (RIR).
|
|
490
|
+
|
|
491
|
+
Args:
|
|
492
|
+
rir_dir: Directory containing .wav RIR files.
|
|
493
|
+
sample_rate: Expected sample rate of audio and RIR files.
|
|
494
|
+
normalize: Normalise output to prevent clipping after convolution.
|
|
495
|
+
p: Probability of applying augmentation.
|
|
496
|
+
"""
|
|
497
|
+
def __init__(
|
|
498
|
+
self,
|
|
499
|
+
rir_dir: str,
|
|
500
|
+
sample_rate: int,
|
|
501
|
+
normalize: bool = True,
|
|
502
|
+
p: float = 0.5,
|
|
503
|
+
):
|
|
504
|
+
self.sample_rate = sample_rate
|
|
505
|
+
self.normalize = normalize
|
|
506
|
+
self.p = p
|
|
507
|
+
|
|
508
|
+
# Pre-load all RIR paths at init — avoids repeated filesystem scans
|
|
509
|
+
self.rir_paths = list(Path(rir_dir).rglob("*.wav"))
|
|
510
|
+
if not self.rir_paths:
|
|
511
|
+
raise ValueError(f"No .wav files found in {rir_dir}")
|
|
512
|
+
|
|
513
|
+
def _load_rir(self, path: Path) -> np.ndarray:
|
|
514
|
+
import soundfile as sf
|
|
515
|
+
rir, sr = sf.read(path, always_2d=True) # (samples, channels)
|
|
516
|
+
rir = rir.T # → (channels, samples)
|
|
517
|
+
if sr != self.sample_rate:
|
|
518
|
+
rir = librosa.resample(rir, orig_sr=sr, target_sr=self.sample_rate)
|
|
519
|
+
return rir.astype(np.float32)
|
|
520
|
+
|
|
521
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
522
|
+
if random.random() > self.p:
|
|
523
|
+
return audio
|
|
524
|
+
|
|
525
|
+
rir_path = random.choice(self.rir_paths)
|
|
526
|
+
rir = self._load_rir(rir_path)
|
|
527
|
+
|
|
528
|
+
audio_np = audio.numpy() # (channels, samples)
|
|
529
|
+
|
|
530
|
+
# Match channel count — mono RIR applied to stereo audio
|
|
531
|
+
if rir.shape[0] == 1 and audio_np.shape[0] > 1:
|
|
532
|
+
rir = np.repeat(rir, audio_np.shape[0], axis=0)
|
|
533
|
+
elif rir.shape[0] > audio_np.shape[0]:
|
|
534
|
+
rir = rir[:audio_np.shape[0]]
|
|
535
|
+
|
|
536
|
+
# Convolve each channel independently
|
|
537
|
+
convolved = np.stack([
|
|
538
|
+
fftconvolve(audio_np[c], rir[c])[:audio_np.shape[1]]
|
|
539
|
+
for c in range(audio_np.shape[0])
|
|
540
|
+
])
|
|
541
|
+
|
|
542
|
+
# Normalise to original peak — convolution can cause large amplitude spikes
|
|
543
|
+
if self.normalize:
|
|
544
|
+
peak = np.abs(convolved).max()
|
|
545
|
+
if peak > 0:
|
|
546
|
+
original_peak = np.abs(audio_np).max()
|
|
547
|
+
convolved = convolved * (original_peak / peak)
|
|
548
|
+
|
|
549
|
+
return torch.from_numpy(convolved).float()
|
|
550
|
+
|
|
551
|
+
def __repr__(self) -> str:
|
|
552
|
+
return (f"{self.__class__.__name__}("
|
|
553
|
+
f"rir_dir='{self.rir_paths[0].parent}', "
|
|
554
|
+
f"n_rirs={len(self.rir_paths)}, p={self.p})")
|
|
555
|
+
|
|
556
|
+
class TimeStretch:
|
|
557
|
+
"""
|
|
558
|
+
Stretch or compress audio in time without changing pitch.
|
|
559
|
+
Uses a phase vocoder via librosa.
|
|
560
|
+
|
|
561
|
+
Args:
|
|
562
|
+
min_rate: Minimum stretch rate (< 1.0 = slower, longer audio).
|
|
563
|
+
max_rate: Maximum stretch rate (> 1.0 = faster, shorter audio).
|
|
564
|
+
p: Probability of applying augmentation.
|
|
565
|
+
"""
|
|
566
|
+
def __init__(
|
|
567
|
+
self,
|
|
568
|
+
min_rate: float = 0.8,
|
|
569
|
+
max_rate: float = 1.25,
|
|
570
|
+
p: float = 0.5,
|
|
571
|
+
):
|
|
572
|
+
if min_rate <= 0 or max_rate <= 0:
|
|
573
|
+
raise ValueError("Rates must be positive")
|
|
574
|
+
if min_rate > max_rate:
|
|
575
|
+
raise ValueError("min_rate must be <= max_rate")
|
|
576
|
+
|
|
577
|
+
self.min_rate = min_rate
|
|
578
|
+
self.max_rate = max_rate
|
|
579
|
+
self.p = p
|
|
580
|
+
|
|
581
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
582
|
+
if random.random() > self.p:
|
|
583
|
+
return audio
|
|
584
|
+
|
|
585
|
+
rate = random.uniform(self.min_rate, self.max_rate)
|
|
586
|
+
|
|
587
|
+
audio_np = audio.numpy() # (channels, samples)
|
|
588
|
+
|
|
589
|
+
# Stretch each channel independently
|
|
590
|
+
stretched = np.stack([
|
|
591
|
+
librosa.effects.time_stretch(audio_np[c], rate=rate)
|
|
592
|
+
for c in range(audio_np.shape[0])
|
|
593
|
+
])
|
|
594
|
+
|
|
595
|
+
return torch.from_numpy(stretched).float()
|
|
596
|
+
|
|
597
|
+
def __repr__(self) -> str:
|
|
598
|
+
return (f"{self.__class__.__name__}("
|
|
599
|
+
f"min_rate={self.min_rate}, max_rate={self.max_rate}, p={self.p})")
|
|
600
|
+
|
|
601
|
+
class ResampleAugmentation:
|
|
602
|
+
"""
|
|
603
|
+
Simulate bandwidth degradation by resampling to an intermediate
|
|
604
|
+
sample rate and back.
|
|
605
|
+
|
|
606
|
+
Args:
|
|
607
|
+
original_freq: Sample rate of the input audio.
|
|
608
|
+
target_freq: Intermediate sample rate to resample to.
|
|
609
|
+
Lower = ResampleDown, Higher = ResampleUp.
|
|
610
|
+
p: Probability of applying augmentation.
|
|
611
|
+
"""
|
|
612
|
+
def __init__(
|
|
613
|
+
self,
|
|
614
|
+
original_freq: int,
|
|
615
|
+
target_freq: int,
|
|
616
|
+
p: float = 0.5,
|
|
617
|
+
):
|
|
618
|
+
self.original_freq = original_freq
|
|
619
|
+
self.target_freq = target_freq
|
|
620
|
+
self.p = p
|
|
621
|
+
|
|
622
|
+
def __call__(self, audio: torch.Tensor) -> torch.Tensor:
|
|
623
|
+
if random.random() > self.p:
|
|
624
|
+
return audio
|
|
625
|
+
|
|
626
|
+
# Step 1: resample to intermediate rate
|
|
627
|
+
degraded = torchaudio.functional.resample(
|
|
628
|
+
audio,
|
|
629
|
+
orig_freq=self.original_freq,
|
|
630
|
+
new_freq=self.target_freq,
|
|
631
|
+
)
|
|
632
|
+
|
|
633
|
+
# Step 2: resample back to original rate
|
|
634
|
+
restored = torchaudio.functional.resample(
|
|
635
|
+
degraded,
|
|
636
|
+
orig_freq=self.target_freq,
|
|
637
|
+
new_freq=self.original_freq,
|
|
638
|
+
)
|
|
639
|
+
|
|
640
|
+
return restored
|
|
641
|
+
|
|
642
|
+
def __repr__(self) -> str:
|
|
643
|
+
direction = "Down" if self.target_freq < self.original_freq else "Up"
|
|
644
|
+
return (
|
|
645
|
+
f"{self.__class__.__name__}(Resample{direction}, "
|
|
646
|
+
f"{self.original_freq}Hz → {self.target_freq}Hz → {self.original_freq}Hz, "
|
|
647
|
+
f"p={self.p})"
|
|
648
|
+
)
|