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,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}")