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.
@@ -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
+ )