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
authtransforms/utils.py
ADDED
|
@@ -0,0 +1,364 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Audio visualization and playback utilities.
|
|
3
|
+
|
|
4
|
+
Works in Jupyter notebooks (IPython display) and plain scripts (matplotlib).
|
|
5
|
+
|
|
6
|
+
Example
|
|
7
|
+
-------
|
|
8
|
+
>>> from utils import plot_audio, play_audio, compare_audio
|
|
9
|
+
>>>
|
|
10
|
+
>>> audio, sr = torchaudio.load('speech.wav')
|
|
11
|
+
>>> augmented = pipeline(audio)
|
|
12
|
+
>>>
|
|
13
|
+
>>> compare_audio(audio, augmented, sr, title_before='Original', title_after='Augmented')
|
|
14
|
+
>>> play_audio(audio, sr, label='Original')
|
|
15
|
+
>>> play_audio(augmented, sr, label='Augmented')
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
import io
|
|
19
|
+
from typing import Optional
|
|
20
|
+
|
|
21
|
+
import matplotlib.pyplot as plt
|
|
22
|
+
import numpy as np
|
|
23
|
+
import torch
|
|
24
|
+
import torchaudio
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
# ---------------------------------------------------------------------------
|
|
28
|
+
# Internal helpers
|
|
29
|
+
# ---------------------------------------------------------------------------
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _to_numpy(audio: torch.Tensor) -> np.ndarray:
|
|
33
|
+
"""Return a 1-D numpy waveform (mono mix if multi-channel)."""
|
|
34
|
+
wav = audio.detach().cpu()
|
|
35
|
+
if wav.dim() == 2 and wav.shape[0] > 1:
|
|
36
|
+
wav = wav.mean(dim=0)
|
|
37
|
+
return wav.squeeze().numpy()
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def _in_notebook() -> bool:
|
|
41
|
+
"""Detect whether we're running inside a Jupyter / IPython kernel."""
|
|
42
|
+
try:
|
|
43
|
+
from IPython import get_ipython
|
|
44
|
+
return get_ipython() is not None
|
|
45
|
+
except ImportError:
|
|
46
|
+
return False
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _tensor_to_wav_bytes(audio: torch.Tensor, sample_rate: int) -> bytes:
|
|
50
|
+
"""Encode a tensor to WAV bytes for IPython playback."""
|
|
51
|
+
buf = io.BytesIO()
|
|
52
|
+
torchaudio.save(buf, audio.cpu(), sample_rate, format="wav")
|
|
53
|
+
buf.seek(0)
|
|
54
|
+
return buf.read()
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
# ---------------------------------------------------------------------------
|
|
58
|
+
# Plotting
|
|
59
|
+
# ---------------------------------------------------------------------------
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def plot_waveform(
|
|
63
|
+
audio: torch.Tensor,
|
|
64
|
+
sample_rate: int,
|
|
65
|
+
title: str = "Waveform",
|
|
66
|
+
ax: Optional[plt.Axes] = None,
|
|
67
|
+
color: str = "#1f77b4",
|
|
68
|
+
) -> plt.Axes:
|
|
69
|
+
"""Plot the raw waveform of an audio tensor.
|
|
70
|
+
|
|
71
|
+
Args:
|
|
72
|
+
audio: Audio tensor of shape (channels, samples) or (samples,).
|
|
73
|
+
sample_rate: Sample rate in Hz.
|
|
74
|
+
title: Plot title.
|
|
75
|
+
ax: Optional existing matplotlib Axes to draw on.
|
|
76
|
+
color: Line color.
|
|
77
|
+
|
|
78
|
+
Returns:
|
|
79
|
+
The Axes object.
|
|
80
|
+
"""
|
|
81
|
+
wav = _to_numpy(audio)
|
|
82
|
+
times = np.arange(len(wav)) / sample_rate
|
|
83
|
+
|
|
84
|
+
if ax is None:
|
|
85
|
+
_, ax = plt.subplots(figsize=(10, 2.5))
|
|
86
|
+
|
|
87
|
+
ax.plot(times, wav, color=color, linewidth=0.6)
|
|
88
|
+
ax.set_title(title)
|
|
89
|
+
ax.set_xlabel("Time (s)")
|
|
90
|
+
ax.set_ylabel("Amplitude")
|
|
91
|
+
ax.set_xlim(times[0], times[-1])
|
|
92
|
+
ax.axhline(0, color="gray", linewidth=0.4, linestyle="--")
|
|
93
|
+
ax.grid(True, alpha=0.3)
|
|
94
|
+
return ax
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def plot_spectrogram(
|
|
98
|
+
audio: torch.Tensor,
|
|
99
|
+
sample_rate: int,
|
|
100
|
+
title: str = "Spectrogram",
|
|
101
|
+
n_fft: int = 512,
|
|
102
|
+
hop_length: int = 128,
|
|
103
|
+
n_mels: Optional[int] = 80,
|
|
104
|
+
ax: Optional[plt.Axes] = None,
|
|
105
|
+
cmap: str = "inferno",
|
|
106
|
+
) -> plt.Axes:
|
|
107
|
+
"""Plot a mel-spectrogram (or linear spectrogram) of an audio tensor.
|
|
108
|
+
|
|
109
|
+
Args:
|
|
110
|
+
audio: Audio tensor of shape (channels, samples) or (samples,).
|
|
111
|
+
sample_rate: Sample rate in Hz.
|
|
112
|
+
title: Plot title.
|
|
113
|
+
n_fft: FFT window size.
|
|
114
|
+
hop_length: STFT hop length.
|
|
115
|
+
n_mels: Number of mel filterbanks. Set to None for a linear spectrogram.
|
|
116
|
+
ax: Optional existing matplotlib Axes.
|
|
117
|
+
cmap: Matplotlib colormap.
|
|
118
|
+
|
|
119
|
+
Returns:
|
|
120
|
+
The Axes object.
|
|
121
|
+
"""
|
|
122
|
+
if n_mels:
|
|
123
|
+
transform = torchaudio.transforms.MelSpectrogram(
|
|
124
|
+
sample_rate=sample_rate,
|
|
125
|
+
n_fft=n_fft,
|
|
126
|
+
hop_length=hop_length,
|
|
127
|
+
n_mels=n_mels,
|
|
128
|
+
)
|
|
129
|
+
else:
|
|
130
|
+
transform = torchaudio.transforms.Spectrogram(n_fft=n_fft, hop_length=hop_length)
|
|
131
|
+
|
|
132
|
+
mono = audio.mean(dim=0, keepdim=True) if audio.shape[0] > 1 else audio
|
|
133
|
+
spec = transform(mono).squeeze(0)
|
|
134
|
+
spec_db = torchaudio.transforms.AmplitudeToDB()(spec).numpy()
|
|
135
|
+
|
|
136
|
+
if ax is None:
|
|
137
|
+
_, ax = plt.subplots(figsize=(10, 3))
|
|
138
|
+
|
|
139
|
+
duration = audio.shape[-1] / sample_rate
|
|
140
|
+
img = ax.imshow(
|
|
141
|
+
spec_db,
|
|
142
|
+
origin="lower",
|
|
143
|
+
aspect="auto",
|
|
144
|
+
extent=[0, duration, 0, sample_rate / 2 / 1000],
|
|
145
|
+
cmap=cmap,
|
|
146
|
+
)
|
|
147
|
+
plt.colorbar(img, ax=ax, format="%+2.0f dB", label="dB")
|
|
148
|
+
ax.set_title(title)
|
|
149
|
+
ax.set_xlabel("Time (s)")
|
|
150
|
+
ax.set_ylabel("Frequency (kHz)")
|
|
151
|
+
return ax
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def plot_audio(
|
|
155
|
+
audio: torch.Tensor,
|
|
156
|
+
sample_rate: int,
|
|
157
|
+
title: str = "Audio",
|
|
158
|
+
n_fft: int = 512,
|
|
159
|
+
hop_length: int = 128,
|
|
160
|
+
n_mels: Optional[int] = 80,
|
|
161
|
+
) -> plt.Figure:
|
|
162
|
+
"""Plot waveform + spectrogram stacked vertically.
|
|
163
|
+
|
|
164
|
+
Args:
|
|
165
|
+
audio: Audio tensor of shape (channels, samples) or (samples,).
|
|
166
|
+
sample_rate: Sample rate in Hz.
|
|
167
|
+
title: Figure suptitle.
|
|
168
|
+
n_fft: FFT window size for spectrogram.
|
|
169
|
+
hop_length: STFT hop length.
|
|
170
|
+
n_mels: Mel filterbanks (None for linear).
|
|
171
|
+
|
|
172
|
+
Returns:
|
|
173
|
+
The Figure.
|
|
174
|
+
"""
|
|
175
|
+
fig, axes = plt.subplots(2, 1, figsize=(12, 5), constrained_layout=True)
|
|
176
|
+
plot_waveform(audio, sample_rate, title="Waveform", ax=axes[0])
|
|
177
|
+
plot_spectrogram(audio, sample_rate, title="Mel-Spectrogram", ax=axes[1],
|
|
178
|
+
n_fft=n_fft, hop_length=hop_length, n_mels=n_mels)
|
|
179
|
+
fig.suptitle(title, fontsize=13, fontweight="bold")
|
|
180
|
+
return fig
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def compare_audio(
|
|
184
|
+
before: torch.Tensor,
|
|
185
|
+
after: torch.Tensor,
|
|
186
|
+
sample_rate: int,
|
|
187
|
+
title_before: str = "Before",
|
|
188
|
+
title_after: str = "After",
|
|
189
|
+
n_fft: int = 512,
|
|
190
|
+
hop_length: int = 128,
|
|
191
|
+
n_mels: Optional[int] = 80,
|
|
192
|
+
) -> plt.Figure:
|
|
193
|
+
"""Side-by-side (before / after) waveform + spectrogram comparison.
|
|
194
|
+
|
|
195
|
+
Plots a 2×2 grid:
|
|
196
|
+
- Row 1: Waveforms (before | after)
|
|
197
|
+
- Row 2: Spectrograms (before | after)
|
|
198
|
+
|
|
199
|
+
Args:
|
|
200
|
+
before: Original audio tensor.
|
|
201
|
+
after: Augmented audio tensor.
|
|
202
|
+
sample_rate: Sample rate in Hz.
|
|
203
|
+
title_before: Column title for the original.
|
|
204
|
+
title_after: Column title for the augmented.
|
|
205
|
+
n_fft: FFT size.
|
|
206
|
+
hop_length: STFT hop length.
|
|
207
|
+
n_mels: Mel filterbanks (None for linear).
|
|
208
|
+
|
|
209
|
+
Returns:
|
|
210
|
+
The Figure.
|
|
211
|
+
"""
|
|
212
|
+
fig, axes = plt.subplots(2, 2, figsize=(16, 6), constrained_layout=True)
|
|
213
|
+
|
|
214
|
+
plot_waveform(before, sample_rate, title=f"{title_before} — Waveform", ax=axes[0, 0], color="#1f77b4")
|
|
215
|
+
plot_waveform(after, sample_rate, title=f"{title_after} — Waveform", ax=axes[0, 1], color="#d62728")
|
|
216
|
+
plot_spectrogram(before, sample_rate, title=f"{title_before} — Spectrogram",
|
|
217
|
+
ax=axes[1, 0], n_fft=n_fft, hop_length=hop_length, n_mels=n_mels)
|
|
218
|
+
plot_spectrogram(after, sample_rate, title=f"{title_after} — Spectrogram",
|
|
219
|
+
ax=axes[1, 1], n_fft=n_fft, hop_length=hop_length, n_mels=n_mels)
|
|
220
|
+
|
|
221
|
+
fig.suptitle("Augmentation Comparison", fontsize=14, fontweight="bold")
|
|
222
|
+
return fig
|
|
223
|
+
|
|
224
|
+
|
|
225
|
+
# ---------------------------------------------------------------------------
|
|
226
|
+
# Audio playback
|
|
227
|
+
# ---------------------------------------------------------------------------
|
|
228
|
+
|
|
229
|
+
|
|
230
|
+
def play_audio(
|
|
231
|
+
audio: torch.Tensor,
|
|
232
|
+
sample_rate: int,
|
|
233
|
+
label: str = "Audio",
|
|
234
|
+
autoplay: bool = False,
|
|
235
|
+
) -> None:
|
|
236
|
+
"""Play audio inline in a Jupyter notebook, or save a temp file and open it.
|
|
237
|
+
|
|
238
|
+
In a Jupyter environment this renders an interactive HTML audio widget.
|
|
239
|
+
In a plain Python script it writes a temporary WAV file and opens the
|
|
240
|
+
system's default media player.
|
|
241
|
+
|
|
242
|
+
Args:
|
|
243
|
+
audio: Audio tensor of shape (channels, samples) or (samples,).
|
|
244
|
+
sample_rate: Sample rate in Hz.
|
|
245
|
+
label: Human-readable label shown above the player (notebook only).
|
|
246
|
+
autoplay: If True, start playback automatically (notebook only).
|
|
247
|
+
"""
|
|
248
|
+
if _in_notebook():
|
|
249
|
+
_play_notebook(audio, sample_rate, label=label, autoplay=autoplay)
|
|
250
|
+
else:
|
|
251
|
+
_play_script(audio, sample_rate)
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
def _play_notebook(
|
|
255
|
+
audio: torch.Tensor,
|
|
256
|
+
sample_rate: int,
|
|
257
|
+
label: str = "Audio",
|
|
258
|
+
autoplay: bool = False,
|
|
259
|
+
) -> None:
|
|
260
|
+
import soundfile as sf
|
|
261
|
+
import tempfile
|
|
262
|
+
import os
|
|
263
|
+
from IPython.display import Audio, display, HTML
|
|
264
|
+
|
|
265
|
+
# write to a temp file — soundfile has no BytesIO issues
|
|
266
|
+
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
|
|
267
|
+
tmp_path = f.name
|
|
268
|
+
|
|
269
|
+
audio_np = audio.numpy()
|
|
270
|
+
if audio_np.ndim == 2:
|
|
271
|
+
audio_np = audio_np.T # (channels, samples) → (samples, channels)
|
|
272
|
+
|
|
273
|
+
sf.write(tmp_path, audio_np, sample_rate)
|
|
274
|
+
display(HTML(f"<p><b>{label}</b></p>"))
|
|
275
|
+
display(Audio(tmp_path, autoplay=autoplay))
|
|
276
|
+
os.unlink(tmp_path)
|
|
277
|
+
|
|
278
|
+
|
|
279
|
+
def _play_script(audio: torch.Tensor, sample_rate: int) -> None:
|
|
280
|
+
"""Write a temp WAV and open it with the OS default player."""
|
|
281
|
+
import platform
|
|
282
|
+
import subprocess
|
|
283
|
+
import tempfile
|
|
284
|
+
|
|
285
|
+
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
|
|
286
|
+
tmp_path = f.name
|
|
287
|
+
|
|
288
|
+
torchaudio.save(tmp_path, audio.cpu(), sample_rate)
|
|
289
|
+
system = platform.system()
|
|
290
|
+
if system == "Darwin":
|
|
291
|
+
subprocess.Popen(["afplay", tmp_path])
|
|
292
|
+
elif system == "Linux":
|
|
293
|
+
subprocess.Popen(["aplay", tmp_path])
|
|
294
|
+
elif system == "Windows":
|
|
295
|
+
import winsound
|
|
296
|
+
winsound.PlaySound(tmp_path, winsound.SND_FILENAME | winsound.SND_ASYNC)
|
|
297
|
+
else:
|
|
298
|
+
print(f"[play_audio] Saved to {tmp_path} — open it manually.")
|
|
299
|
+
|
|
300
|
+
|
|
301
|
+
def compare_play(
|
|
302
|
+
before: torch.Tensor,
|
|
303
|
+
after: torch.Tensor,
|
|
304
|
+
sample_rate: int,
|
|
305
|
+
label_before: str = "Before",
|
|
306
|
+
label_after: str = "After",
|
|
307
|
+
) -> None:
|
|
308
|
+
"""Play both before and after audio widgets (notebook) or sequentially (script).
|
|
309
|
+
|
|
310
|
+
Args:
|
|
311
|
+
before: Original audio tensor.
|
|
312
|
+
after: Augmented audio tensor.
|
|
313
|
+
sample_rate: Sample rate.
|
|
314
|
+
label_before: Label for the original.
|
|
315
|
+
label_after: Label for the augmented.
|
|
316
|
+
"""
|
|
317
|
+
play_audio(before, sample_rate, label=label_before)
|
|
318
|
+
play_audio(after, sample_rate, label=label_after)
|
|
319
|
+
|
|
320
|
+
|
|
321
|
+
# ---------------------------------------------------------------------------
|
|
322
|
+
# Quick info helper
|
|
323
|
+
# ---------------------------------------------------------------------------
|
|
324
|
+
|
|
325
|
+
|
|
326
|
+
def audio_info(audio: torch.Tensor, sample_rate: int, label: str = "Audio") -> None:
|
|
327
|
+
"""Print a summary of an audio tensor's properties.
|
|
328
|
+
|
|
329
|
+
Args:
|
|
330
|
+
audio: Audio tensor.
|
|
331
|
+
sample_rate: Sample rate in Hz.
|
|
332
|
+
label: Human-readable name shown in the output.
|
|
333
|
+
"""
|
|
334
|
+
channels = audio.shape[0] if audio.dim() == 2 else 1
|
|
335
|
+
samples = audio.shape[-1]
|
|
336
|
+
duration = samples / sample_rate
|
|
337
|
+
peak = audio.abs().max().item()
|
|
338
|
+
rms = audio.pow(2).mean().sqrt().item()
|
|
339
|
+
rms_db = 20 * np.log10(rms + 1e-9)
|
|
340
|
+
|
|
341
|
+
print(f"[{label}]")
|
|
342
|
+
print(f" Shape : {tuple(audio.shape)}")
|
|
343
|
+
print(f" Channels : {channels}")
|
|
344
|
+
print(f" Sample rate: {sample_rate} Hz")
|
|
345
|
+
print(f" Duration : {duration:.3f} s ({samples} samples)")
|
|
346
|
+
print(f" Peak : {peak:.4f}")
|
|
347
|
+
print(f" RMS : {rms:.4f} ({rms_db:.1f} dBFS)")
|
|
348
|
+
print(f" dtype : {audio.dtype}")
|
|
349
|
+
|
|
350
|
+
def save_audio(audio: torch.Tensor, sample_rate: int, path: str = "output.wav") -> str:
|
|
351
|
+
"""
|
|
352
|
+
Save a torch audio tensor to a .wav file.
|
|
353
|
+
|
|
354
|
+
Args:
|
|
355
|
+
audio: Audio tensor of shape (channels, samples)
|
|
356
|
+
sample_rate: Sample rate of the audio
|
|
357
|
+
path: Output file path (default: output.wav)
|
|
358
|
+
|
|
359
|
+
Returns:
|
|
360
|
+
Absolute path of the saved file
|
|
361
|
+
"""
|
|
362
|
+
import soundfile as sf
|
|
363
|
+
sf.write(path, audio.numpy().T, sample_rate)
|
|
364
|
+
print(f"Saved → {path}")
|