commkit 1.0.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.
- commkit/__init__.py +74 -0
- commkit/_cuda/__init__.py +321 -0
- commkit/_cuda/compiler.py +88 -0
- commkit/_cuda/src/bps_min_d2.cu +104 -0
- commkit/_cuda/src/cs_block.cu +119 -0
- commkit/_cuda/src/selftest.cu +14 -0
- commkit/analysis/__init__.py +55 -0
- commkit/analysis/_common.py +236 -0
- commkit/analysis/allan.py +108 -0
- commkit/analysis/drift.py +213 -0
- commkit/analysis/interferometry.py +887 -0
- commkit/analysis/linewidth.py +480 -0
- commkit/analysis/trajectory.py +91 -0
- commkit/backend.py +507 -0
- commkit/coding/__init__.py +23 -0
- commkit/coding/base.py +17 -0
- commkit/coding/bch.py +6 -0
- commkit/coding/convolutional.py +7 -0
- commkit/coding/crc.py +7 -0
- commkit/coding/galois.py +8 -0
- commkit/coding/hamming.py +6 -0
- commkit/coding/interleaving.py +7 -0
- commkit/coding/ldpc.py +8 -0
- commkit/coding/polar.py +8 -0
- commkit/coding/ratematch.py +6 -0
- commkit/coding/reed_solomon.py +6 -0
- commkit/coding/turbo.py +8 -0
- commkit/core/__init__.py +32 -0
- commkit/core/frame.py +992 -0
- commkit/core/generation.py +581 -0
- commkit/core/signal.py +725 -0
- commkit/equalization/__init__.py +49 -0
- commkit/equalization/_block.py +1855 -0
- commkit/equalization/_common.py +606 -0
- commkit/equalization/_kernels_jax.py +1720 -0
- commkit/equalization/_kernels_numba.py +1704 -0
- commkit/equalization/blind.py +223 -0
- commkit/equalization/linear.py +365 -0
- commkit/equalization/polarization.py +790 -0
- commkit/equalization/result.py +191 -0
- commkit/equalization/sequential.py +2805 -0
- commkit/filtering.py +1120 -0
- commkit/frequency.py +1191 -0
- commkit/helpers.py +489 -0
- commkit/impairments/__init__.py +43 -0
- commkit/impairments/channel/__init__.py +20 -0
- commkit/impairments/channel/linear.py +310 -0
- commkit/impairments/channel/nonlinear.py +11 -0
- commkit/impairments/frontend.py +229 -0
- commkit/impairments/noise.py +105 -0
- commkit/impairments/source.py +219 -0
- commkit/io.py +308 -0
- commkit/logger.py +103 -0
- commkit/mapping/__init__.py +46 -0
- commkit/mapping/bits.py +240 -0
- commkit/mapping/constellation.py +153 -0
- commkit/mapping/gray.py +429 -0
- commkit/mapping/llr.py +253 -0
- commkit/mapping/shaping.py +218 -0
- commkit/metrics.py +949 -0
- commkit/multirate.py +476 -0
- commkit/plotting/__init__.py +78 -0
- commkit/plotting/analysis.py +627 -0
- commkit/plotting/constellation.py +483 -0
- commkit/plotting/equalizer.py +390 -0
- commkit/plotting/eye.py +388 -0
- commkit/plotting/spectral.py +575 -0
- commkit/plotting/sync.py +953 -0
- commkit/plotting/theme.py +203 -0
- commkit/plotting/waveform.py +200 -0
- commkit/py.typed +0 -0
- commkit/recovery/__init__.py +51 -0
- commkit/recovery/bps.py +337 -0
- commkit/recovery/corrections.py +751 -0
- commkit/recovery/pilots.py +803 -0
- commkit/recovery/pll.py +482 -0
- commkit/recovery/tikhonov.py +424 -0
- commkit/recovery/viterbi_viterbi.py +227 -0
- commkit/spectral.py +560 -0
- commkit/timing.py +841 -0
- commkit-1.0.0.dist-info/METADATA +145 -0
- commkit-1.0.0.dist-info/RECORD +84 -0
- commkit-1.0.0.dist-info/WHEEL +4 -0
- commkit-1.0.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,575 @@
|
|
|
1
|
+
"""Power spectral density and spectrogram plots."""
|
|
2
|
+
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import matplotlib.pyplot as plt
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from ..backend import dispatch, to_device
|
|
9
|
+
from ..core.signal import Signal
|
|
10
|
+
from ..logger import logger
|
|
11
|
+
from .theme import (
|
|
12
|
+
_create_subplot_grid,
|
|
13
|
+
_grid_figsize,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def plot_psd(
|
|
18
|
+
samples: Any,
|
|
19
|
+
sampling_rate: float = 1.0,
|
|
20
|
+
nperseg: int = 256,
|
|
21
|
+
detrend: str | bool | None = False,
|
|
22
|
+
average: str | None = "mean",
|
|
23
|
+
window: str | tuple[Any, ...] | Any = "hann",
|
|
24
|
+
noverlap: int | None = None,
|
|
25
|
+
nfft: int | None = None,
|
|
26
|
+
scaling: str = "density",
|
|
27
|
+
center_frequency: float = 0.0,
|
|
28
|
+
domain: str = "RF",
|
|
29
|
+
x_axis: str = "frequency",
|
|
30
|
+
ax: Any | None = None,
|
|
31
|
+
xlim: tuple[float, float] | None = None,
|
|
32
|
+
ylim: tuple[float, float] | None = None,
|
|
33
|
+
title: str | None = "Power Spectral Density",
|
|
34
|
+
show: bool = False,
|
|
35
|
+
**kwargs: Any,
|
|
36
|
+
) -> tuple[Any, Any] | None:
|
|
37
|
+
"""
|
|
38
|
+
Plots the Power Spectral Density (PSD) of the signal.
|
|
39
|
+
|
|
40
|
+
Supports automatic frequency scaling (Hz, MHz, GHz, etc.) or wavelength
|
|
41
|
+
conversion for optical signals. Handles multidimensional (MIMO) signals
|
|
42
|
+
by generating a grid of subplots.
|
|
43
|
+
|
|
44
|
+
Parameters
|
|
45
|
+
----------
|
|
46
|
+
samples : array_like or Signal
|
|
47
|
+
Input signal samples. Shape: (..., N_samples).
|
|
48
|
+
sampling_rate : float, default 1.0
|
|
49
|
+
Sampling rate in Hz.
|
|
50
|
+
nperseg : int, default 256
|
|
51
|
+
Length of each segment for Welch's method. Higher values provide
|
|
52
|
+
better frequency resolution but more noise.
|
|
53
|
+
detrend : str or bool, default False
|
|
54
|
+
Specifies how to detrend each segment (e.g., 'constant', 'linear').
|
|
55
|
+
average : str, default "mean"
|
|
56
|
+
Method to use for averaging segments ('mean' or 'median').
|
|
57
|
+
window : str or tuple or array_like, default "hann"
|
|
58
|
+
Desired window to use. If `window` is a string or tuple, it is
|
|
59
|
+
passed to `scipy.signal.get_window` to generate the window values.
|
|
60
|
+
noverlap : int, optional
|
|
61
|
+
Number of points to overlap between segments. If None,
|
|
62
|
+
`noverlap = nperseg // 2`.
|
|
63
|
+
nfft : int, optional
|
|
64
|
+
Length of the FFT used, if a zero padded FFT is desired. If None,
|
|
65
|
+
the FFT length is `nperseg`.
|
|
66
|
+
scaling : {"density", "spectrum"}, default "density"
|
|
67
|
+
Selects between computing the power spectral density ('density')
|
|
68
|
+
where Pxx has units of V**2/Hz and computing the power spectrum
|
|
69
|
+
('spectrum') where Pxx has units of V**2.
|
|
70
|
+
center_frequency : float, default 0.0
|
|
71
|
+
Frequency offset to apply to the x-axis in Hz.
|
|
72
|
+
domain : {"RF", "OPT"}, default "RF"
|
|
73
|
+
Signal domain. If "OPT", wavelength scaling is enabled.
|
|
74
|
+
x_axis : {"frequency", "wavelength"}, default "frequency"
|
|
75
|
+
Units for the horizontal axis.
|
|
76
|
+
ax : matplotlib.axes.Axes, optional
|
|
77
|
+
Existing axis to plot on. If `None`, a new figure is created.
|
|
78
|
+
xlim, ylim : tuple of float, optional
|
|
79
|
+
Axis limits for the plot.
|
|
80
|
+
title : str, optional
|
|
81
|
+
Plot title. Defaults to "Power Spectral Density".
|
|
82
|
+
show : bool, default False
|
|
83
|
+
If True, calls `plt.show()` immediately.
|
|
84
|
+
**kwargs : Any
|
|
85
|
+
Additional keyword arguments passed to `ax.plot`.
|
|
86
|
+
|
|
87
|
+
Returns
|
|
88
|
+
-------
|
|
89
|
+
fig : matplotlib.figure.Figure
|
|
90
|
+
The figure object.
|
|
91
|
+
ax : matplotlib.axes.Axes or ndarray
|
|
92
|
+
The axis or array of axes used for the plot.
|
|
93
|
+
"""
|
|
94
|
+
if isinstance(samples, Signal):
|
|
95
|
+
sig = samples
|
|
96
|
+
return plot_psd(
|
|
97
|
+
sig.samples,
|
|
98
|
+
sampling_rate=sig.sampling_rate,
|
|
99
|
+
nperseg=nperseg,
|
|
100
|
+
detrend=detrend,
|
|
101
|
+
average=average,
|
|
102
|
+
window=window,
|
|
103
|
+
noverlap=noverlap,
|
|
104
|
+
nfft=nfft,
|
|
105
|
+
scaling=scaling,
|
|
106
|
+
center_frequency=sig.center_frequency,
|
|
107
|
+
domain=sig.physical_domain or "RF",
|
|
108
|
+
x_axis=x_axis,
|
|
109
|
+
ax=ax,
|
|
110
|
+
xlim=xlim,
|
|
111
|
+
ylim=ylim,
|
|
112
|
+
title=title,
|
|
113
|
+
show=show,
|
|
114
|
+
**kwargs,
|
|
115
|
+
)
|
|
116
|
+
|
|
117
|
+
logger.debug("Generating PSD plot (sampling_rate=%s Hz).", sampling_rate)
|
|
118
|
+
|
|
119
|
+
samples, xp, _ = dispatch(samples)
|
|
120
|
+
|
|
121
|
+
# Handle Multichannel (e.g. Dual-Pol)
|
|
122
|
+
# Convention: (Channels, Time)
|
|
123
|
+
if samples.ndim > 1:
|
|
124
|
+
num_channels = samples.shape[0]
|
|
125
|
+
|
|
126
|
+
if ax is None:
|
|
127
|
+
nrows, ncols = _create_subplot_grid(num_channels)
|
|
128
|
+
fig, axes = plt.subplots(
|
|
129
|
+
nrows, ncols, figsize=_grid_figsize(nrows, ncols), squeeze=False
|
|
130
|
+
)
|
|
131
|
+
else:
|
|
132
|
+
if not isinstance(ax, (list, tuple, np.ndarray)):
|
|
133
|
+
# If single axis provided but multiple channels, warn and overlay?
|
|
134
|
+
# Or better, just overlay on the same axis for PSD
|
|
135
|
+
# Actually, the user asked for "side-by-side".
|
|
136
|
+
# But if the user provides a single axis, we must respect it.
|
|
137
|
+
# Let's overlay if single axis provided, or fail.
|
|
138
|
+
logger.warning(
|
|
139
|
+
"Multiple channels detected but single axis provided. Overlaying plots."
|
|
140
|
+
)
|
|
141
|
+
axes = np.array([[ax] * num_channels])
|
|
142
|
+
fig = ax.figure
|
|
143
|
+
else:
|
|
144
|
+
axes = np.atleast_2d(ax)
|
|
145
|
+
fig = axes.flat[0].figure
|
|
146
|
+
|
|
147
|
+
for i in range(num_channels):
|
|
148
|
+
# Recursively call psd for each channel
|
|
149
|
+
channel_samples = samples[i]
|
|
150
|
+
|
|
151
|
+
# Determine target axis using 2D indexing
|
|
152
|
+
row, col = divmod(i, axes.shape[1])
|
|
153
|
+
target_ax = axes[row, col] if row < axes.shape[0] else axes.flat[-1]
|
|
154
|
+
|
|
155
|
+
ch_title = f"{title} (Ch {i})" if title else f"Channel {i}"
|
|
156
|
+
|
|
157
|
+
plot_psd(
|
|
158
|
+
channel_samples,
|
|
159
|
+
sampling_rate=sampling_rate,
|
|
160
|
+
nperseg=nperseg,
|
|
161
|
+
detrend=detrend,
|
|
162
|
+
average=average,
|
|
163
|
+
window=window,
|
|
164
|
+
noverlap=noverlap,
|
|
165
|
+
nfft=nfft,
|
|
166
|
+
scaling=scaling,
|
|
167
|
+
center_frequency=center_frequency,
|
|
168
|
+
domain=domain,
|
|
169
|
+
x_axis=x_axis,
|
|
170
|
+
ax=target_ax,
|
|
171
|
+
xlim=xlim,
|
|
172
|
+
ylim=ylim,
|
|
173
|
+
title=ch_title,
|
|
174
|
+
show=False,
|
|
175
|
+
**kwargs,
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
if show:
|
|
179
|
+
plt.show()
|
|
180
|
+
return None
|
|
181
|
+
return fig, axes
|
|
182
|
+
|
|
183
|
+
# --- 1D Logic Starts Here ---
|
|
184
|
+
|
|
185
|
+
if ax is None:
|
|
186
|
+
fig, ax = plt.subplots()
|
|
187
|
+
else:
|
|
188
|
+
fig = ax.figure
|
|
189
|
+
|
|
190
|
+
from .. import spectral
|
|
191
|
+
|
|
192
|
+
# Calculate PSD
|
|
193
|
+
f, Pxx = spectral.welch_psd(
|
|
194
|
+
samples,
|
|
195
|
+
sampling_rate=sampling_rate,
|
|
196
|
+
nperseg=nperseg,
|
|
197
|
+
detrend=detrend,
|
|
198
|
+
average=average,
|
|
199
|
+
window=window,
|
|
200
|
+
noverlap=noverlap,
|
|
201
|
+
nfft=nfft,
|
|
202
|
+
scaling=scaling,
|
|
203
|
+
)
|
|
204
|
+
|
|
205
|
+
# Move to cpu for plotting
|
|
206
|
+
f = to_device(f, "cpu")
|
|
207
|
+
Pxx = to_device(Pxx, "cpu")
|
|
208
|
+
|
|
209
|
+
# Apply center frequency shift
|
|
210
|
+
f = f + center_frequency
|
|
211
|
+
|
|
212
|
+
xlabel = "Frequency [Hz]"
|
|
213
|
+
x_values = f
|
|
214
|
+
|
|
215
|
+
if x_axis == "wavelength":
|
|
216
|
+
if domain != "OPT":
|
|
217
|
+
logger.warning("Wavelength plotting is typically used for optical signals.")
|
|
218
|
+
|
|
219
|
+
# c = 299,792,458 m/s
|
|
220
|
+
c = 299792458.0
|
|
221
|
+
# Avoid division by zero
|
|
222
|
+
# Convert frequency to wavelength: lambda = c / f
|
|
223
|
+
# Result in nanometers (1e9)
|
|
224
|
+
valid_indices = f > 0
|
|
225
|
+
x_values = np.zeros_like(f)
|
|
226
|
+
x_values[valid_indices] = (c / f[valid_indices]) * 1e9
|
|
227
|
+
x_values[~valid_indices] = np.nan # Handle non-positive frequencies
|
|
228
|
+
|
|
229
|
+
xlabel = "Wavelength [nm]"
|
|
230
|
+
else:
|
|
231
|
+
# Auto-scale frequency axis
|
|
232
|
+
max_f = np.max(np.abs(f))
|
|
233
|
+
if max_f >= 1e12:
|
|
234
|
+
scale_factor = 1e12
|
|
235
|
+
unit = "THz"
|
|
236
|
+
elif max_f >= 1e9:
|
|
237
|
+
scale_factor = 1e9
|
|
238
|
+
unit = "GHz"
|
|
239
|
+
elif max_f >= 1e6:
|
|
240
|
+
scale_factor = 1e6
|
|
241
|
+
unit = "MHz"
|
|
242
|
+
elif max_f >= 1e3:
|
|
243
|
+
scale_factor = 1e3
|
|
244
|
+
unit = "kHz"
|
|
245
|
+
else:
|
|
246
|
+
scale_factor = 1.0
|
|
247
|
+
unit = "Hz"
|
|
248
|
+
|
|
249
|
+
x_values = f / scale_factor
|
|
250
|
+
xlabel = f"Frequency [{unit}]"
|
|
251
|
+
|
|
252
|
+
if xlim is not None:
|
|
253
|
+
ax.set_xlim(xlim)
|
|
254
|
+
if ylim is not None:
|
|
255
|
+
ax.set_ylim(ylim)
|
|
256
|
+
|
|
257
|
+
# Add epsilon to avoid log(0) warnings
|
|
258
|
+
ax.plot(x_values, 10 * np.log10(Pxx + 1e-20), **kwargs)
|
|
259
|
+
ax.set_xlabel(xlabel)
|
|
260
|
+
ax.set_ylabel("PSD [dB/Hz]")
|
|
261
|
+
if title is not None:
|
|
262
|
+
ax.set_title(title)
|
|
263
|
+
|
|
264
|
+
if show:
|
|
265
|
+
plt.show()
|
|
266
|
+
return None
|
|
267
|
+
return fig, ax
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
def plot_spectrogram(
|
|
271
|
+
samples: Any,
|
|
272
|
+
sampling_rate: float = 1.0,
|
|
273
|
+
window: str | tuple[Any, ...] | Any = "hann",
|
|
274
|
+
nperseg: int = 256,
|
|
275
|
+
noverlap: int | None = None,
|
|
276
|
+
nfft: int | None = None,
|
|
277
|
+
detrend: str | bool | None = False,
|
|
278
|
+
return_onesided: bool | None = None,
|
|
279
|
+
scaling: str = "density",
|
|
280
|
+
axis: int = -1,
|
|
281
|
+
mode: str = "psd",
|
|
282
|
+
center_frequency: float = 0.0,
|
|
283
|
+
domain: str = "RF",
|
|
284
|
+
ax: Any | None = None,
|
|
285
|
+
xlim: tuple[float, float] | None = None,
|
|
286
|
+
ylim: tuple[float, float] | None = None,
|
|
287
|
+
title: str | None = "Spectrogram",
|
|
288
|
+
cmap: str = "viridis",
|
|
289
|
+
show: bool = False,
|
|
290
|
+
**kwargs: Any,
|
|
291
|
+
) -> tuple[Any, Any] | None:
|
|
292
|
+
"""
|
|
293
|
+
Plots the spectrogram of a signal.
|
|
294
|
+
|
|
295
|
+
Plots with Frequency on the horizontal axis (x-axis) and Time on the
|
|
296
|
+
vertical axis (y-axis). Supports dynamic subplots for multichannel/MIMO
|
|
297
|
+
signals.
|
|
298
|
+
|
|
299
|
+
Parameters
|
|
300
|
+
----------
|
|
301
|
+
samples : array_like or Signal
|
|
302
|
+
Input signal samples. Shape: (..., N_samples).
|
|
303
|
+
sampling_rate : float, default 1.0
|
|
304
|
+
Sampling rate in Hz.
|
|
305
|
+
window : str or tuple or array_like, default "hann"
|
|
306
|
+
Desired window to use.
|
|
307
|
+
nperseg : int, default 256
|
|
308
|
+
Length of each segment.
|
|
309
|
+
noverlap : int, optional
|
|
310
|
+
Number of points to overlap between segments.
|
|
311
|
+
nfft : int, optional
|
|
312
|
+
Length of the FFT used.
|
|
313
|
+
detrend : str or bool, default False
|
|
314
|
+
Specifies how to detrend each segment.
|
|
315
|
+
return_onesided : bool, optional
|
|
316
|
+
If True, returns a one-sided spectrum for real-valued data.
|
|
317
|
+
scaling : {"density", "spectrum"}, default "density"
|
|
318
|
+
Selects between computing power spectral density or power spectrum.
|
|
319
|
+
axis : int, default -1
|
|
320
|
+
The axis along which to compute the spectrogram.
|
|
321
|
+
mode : {"psd", "complex", "magnitude", "angle", "phase"}, default "psd"
|
|
322
|
+
Type of spectrogram to return.
|
|
323
|
+
center_frequency : float, default 0.0
|
|
324
|
+
Frequency offset to apply to the frequency axis in Hz.
|
|
325
|
+
domain : {"RF", "OPT"}, default "RF"
|
|
326
|
+
Signal domain.
|
|
327
|
+
ax : matplotlib.axes.Axes, optional
|
|
328
|
+
Existing axis to plot on.
|
|
329
|
+
xlim : tuple of float, optional
|
|
330
|
+
Frequency limits for the plot (in Hz, after center_frequency offset).
|
|
331
|
+
Used to crop data before plotting for performance.
|
|
332
|
+
ylim : tuple of float, optional
|
|
333
|
+
Time limits for the plot (in seconds).
|
|
334
|
+
Used to crop data before plotting for performance.
|
|
335
|
+
title : str, optional
|
|
336
|
+
Plot title.
|
|
337
|
+
cmap : str, default "viridis"
|
|
338
|
+
Colormap for the spectrogram plot.
|
|
339
|
+
show : bool, default False
|
|
340
|
+
If True, calls `plt.show()` immediately.
|
|
341
|
+
**kwargs : Any
|
|
342
|
+
Additional keyword arguments passed to `ax.pcolormesh`.
|
|
343
|
+
|
|
344
|
+
Returns
|
|
345
|
+
-------
|
|
346
|
+
fig : matplotlib.figure.Figure
|
|
347
|
+
The figure object.
|
|
348
|
+
ax : matplotlib.axes.Axes or ndarray
|
|
349
|
+
The axis or array of axes used for the plot.
|
|
350
|
+
"""
|
|
351
|
+
if isinstance(samples, Signal):
|
|
352
|
+
sig = samples
|
|
353
|
+
return plot_spectrogram(
|
|
354
|
+
sig.samples,
|
|
355
|
+
sampling_rate=sig.sampling_rate,
|
|
356
|
+
window=window,
|
|
357
|
+
nperseg=nperseg,
|
|
358
|
+
noverlap=noverlap,
|
|
359
|
+
nfft=nfft,
|
|
360
|
+
detrend=detrend,
|
|
361
|
+
return_onesided=return_onesided,
|
|
362
|
+
scaling=scaling,
|
|
363
|
+
axis=-1,
|
|
364
|
+
mode=mode,
|
|
365
|
+
center_frequency=sig.center_frequency,
|
|
366
|
+
domain=sig.physical_domain or "RF",
|
|
367
|
+
ax=ax,
|
|
368
|
+
xlim=xlim,
|
|
369
|
+
ylim=ylim,
|
|
370
|
+
title=title,
|
|
371
|
+
cmap=cmap,
|
|
372
|
+
show=show,
|
|
373
|
+
**kwargs,
|
|
374
|
+
)
|
|
375
|
+
|
|
376
|
+
logger.debug("Generating spectrogram plot (sampling_rate=%s Hz).", sampling_rate)
|
|
377
|
+
|
|
378
|
+
samples, xp, _ = dispatch(samples)
|
|
379
|
+
|
|
380
|
+
# Handle Multichannel (MIMO)
|
|
381
|
+
# Convention: (Channels, Time)
|
|
382
|
+
if samples.ndim > 1:
|
|
383
|
+
num_channels = samples.shape[0]
|
|
384
|
+
|
|
385
|
+
if ax is None:
|
|
386
|
+
nrows, ncols = _create_subplot_grid(num_channels)
|
|
387
|
+
fig, axes = plt.subplots(
|
|
388
|
+
nrows, ncols, figsize=_grid_figsize(nrows, ncols), squeeze=False
|
|
389
|
+
)
|
|
390
|
+
else:
|
|
391
|
+
if not isinstance(ax, (list, tuple, np.ndarray)):
|
|
392
|
+
logger.warning(
|
|
393
|
+
"Multiple channels detected but single axis provided. Overlaying plots."
|
|
394
|
+
)
|
|
395
|
+
axes = np.array([[ax] * num_channels])
|
|
396
|
+
fig = ax.figure
|
|
397
|
+
else:
|
|
398
|
+
axes = np.atleast_2d(ax)
|
|
399
|
+
fig = axes.flat[0].figure
|
|
400
|
+
|
|
401
|
+
for i in range(num_channels):
|
|
402
|
+
channel_samples = samples[i]
|
|
403
|
+
row, col = divmod(i, axes.shape[1])
|
|
404
|
+
target_ax = axes[row, col] if row < axes.shape[0] else axes.flat[-1]
|
|
405
|
+
ch_title = f"{title} (Ch {i})" if title else f"Channel {i}"
|
|
406
|
+
|
|
407
|
+
plot_spectrogram(
|
|
408
|
+
channel_samples,
|
|
409
|
+
sampling_rate=sampling_rate,
|
|
410
|
+
window=window,
|
|
411
|
+
nperseg=nperseg,
|
|
412
|
+
noverlap=noverlap,
|
|
413
|
+
nfft=nfft,
|
|
414
|
+
detrend=detrend,
|
|
415
|
+
return_onesided=return_onesided,
|
|
416
|
+
scaling=scaling,
|
|
417
|
+
axis=axis,
|
|
418
|
+
mode=mode,
|
|
419
|
+
center_frequency=center_frequency,
|
|
420
|
+
domain=domain,
|
|
421
|
+
ax=target_ax,
|
|
422
|
+
xlim=xlim,
|
|
423
|
+
ylim=ylim,
|
|
424
|
+
title=ch_title,
|
|
425
|
+
cmap=cmap,
|
|
426
|
+
show=False,
|
|
427
|
+
**kwargs,
|
|
428
|
+
)
|
|
429
|
+
|
|
430
|
+
if show:
|
|
431
|
+
plt.show()
|
|
432
|
+
return None
|
|
433
|
+
return fig, axes
|
|
434
|
+
|
|
435
|
+
# --- 1D Logic ---
|
|
436
|
+
if ax is None:
|
|
437
|
+
fig, ax = plt.subplots()
|
|
438
|
+
else:
|
|
439
|
+
fig = ax.figure
|
|
440
|
+
|
|
441
|
+
from .. import spectral
|
|
442
|
+
|
|
443
|
+
# Calculate spectrogram
|
|
444
|
+
f, t, Sxx = spectral.spectrogram(
|
|
445
|
+
samples,
|
|
446
|
+
sampling_rate=sampling_rate,
|
|
447
|
+
window=window,
|
|
448
|
+
nperseg=nperseg,
|
|
449
|
+
noverlap=noverlap,
|
|
450
|
+
nfft=nfft,
|
|
451
|
+
detrend=detrend,
|
|
452
|
+
return_onesided=return_onesided,
|
|
453
|
+
scaling=scaling,
|
|
454
|
+
axis=axis,
|
|
455
|
+
mode=mode,
|
|
456
|
+
)
|
|
457
|
+
|
|
458
|
+
# Move to CPU for plotting
|
|
459
|
+
f = to_device(f, "cpu")
|
|
460
|
+
t = to_device(t, "cpu")
|
|
461
|
+
Sxx = to_device(Sxx, "cpu")
|
|
462
|
+
|
|
463
|
+
# Shift frequency axis first
|
|
464
|
+
f_shifted = f + center_frequency
|
|
465
|
+
|
|
466
|
+
# Masking frequency axis (xlim corresponds to frequency axis)
|
|
467
|
+
f_start, f_end = 0, len(f_shifted)
|
|
468
|
+
if xlim is not None:
|
|
469
|
+
f_mask = (f_shifted >= xlim[0]) & (f_shifted <= xlim[1])
|
|
470
|
+
f_indices = np.where(f_mask)[0]
|
|
471
|
+
if len(f_indices) > 0:
|
|
472
|
+
f_start, f_end = f_indices[0], f_indices[-1] + 1
|
|
473
|
+
else:
|
|
474
|
+
logger.warning(
|
|
475
|
+
"xlim (frequency) %s does not overlap with frequency range [%.3f, %.3f]. Plotting whole frequency range.",
|
|
476
|
+
xlim,
|
|
477
|
+
f_shifted[0],
|
|
478
|
+
f_shifted[-1],
|
|
479
|
+
)
|
|
480
|
+
|
|
481
|
+
# Masking time axis (ylim corresponds to time axis)
|
|
482
|
+
t_start, t_end = 0, len(t)
|
|
483
|
+
if ylim is not None:
|
|
484
|
+
t_mask = (t >= ylim[0]) & (t <= ylim[1])
|
|
485
|
+
t_indices = np.where(t_mask)[0]
|
|
486
|
+
if len(t_indices) > 0:
|
|
487
|
+
t_start, t_end = t_indices[0], t_indices[-1] + 1
|
|
488
|
+
else:
|
|
489
|
+
logger.warning(
|
|
490
|
+
"ylim (time) %s does not overlap with time range [%.3f, %.3f]. Plotting whole time axis.",
|
|
491
|
+
ylim,
|
|
492
|
+
t[0],
|
|
493
|
+
t[-1],
|
|
494
|
+
)
|
|
495
|
+
|
|
496
|
+
# Slice arrays for plotting performance
|
|
497
|
+
f_plot = f_shifted[f_start:f_end]
|
|
498
|
+
t_plot = t[t_start:t_end]
|
|
499
|
+
Sxx_slice = Sxx[f_start:f_end, t_start:t_end]
|
|
500
|
+
|
|
501
|
+
# Convert values based on mode (e.g. dB scale for PSD/magnitude)
|
|
502
|
+
if mode == "psd":
|
|
503
|
+
Sxx_plot = 10 * np.log10(Sxx_slice + 1e-20)
|
|
504
|
+
elif mode in ("complex", "magnitude"):
|
|
505
|
+
Sxx_plot = 10 * np.log10(np.abs(Sxx_slice) ** 2 + 1e-20)
|
|
506
|
+
else:
|
|
507
|
+
# Angle, phase, etc., plot linearly
|
|
508
|
+
Sxx_plot = Sxx_slice
|
|
509
|
+
|
|
510
|
+
# Auto-scale frequency axis (x-axis)
|
|
511
|
+
max_f = np.max(np.abs(f_plot)) if len(f_plot) > 0 else 0
|
|
512
|
+
if max_f >= 1e12:
|
|
513
|
+
f_scale = 1e12
|
|
514
|
+
f_unit = "THz"
|
|
515
|
+
elif max_f >= 1e9:
|
|
516
|
+
f_scale = 1e9
|
|
517
|
+
f_unit = "GHz"
|
|
518
|
+
elif max_f >= 1e6:
|
|
519
|
+
f_scale = 1e6
|
|
520
|
+
f_unit = "MHz"
|
|
521
|
+
elif max_f >= 1e3:
|
|
522
|
+
f_scale = 1e3
|
|
523
|
+
f_unit = "kHz"
|
|
524
|
+
else:
|
|
525
|
+
f_scale = 1.0
|
|
526
|
+
f_unit = "Hz"
|
|
527
|
+
|
|
528
|
+
x_values = f_plot / f_scale
|
|
529
|
+
xlabel = f"Frequency [{f_unit}]"
|
|
530
|
+
|
|
531
|
+
# Auto-scale time axis (y-axis)
|
|
532
|
+
max_t = t_plot[-1] if len(t_plot) > 0 else 0
|
|
533
|
+
if max_t < 1e-9:
|
|
534
|
+
t_scale = 1e12
|
|
535
|
+
t_unit = "ps"
|
|
536
|
+
elif max_t < 1e-6:
|
|
537
|
+
t_scale = 1e9
|
|
538
|
+
t_unit = "ns"
|
|
539
|
+
elif max_t < 1e-3:
|
|
540
|
+
t_scale = 1e6
|
|
541
|
+
t_unit = "µs"
|
|
542
|
+
elif max_t < 1:
|
|
543
|
+
t_scale = 1e3
|
|
544
|
+
t_unit = "ms"
|
|
545
|
+
else:
|
|
546
|
+
t_scale = 1.0
|
|
547
|
+
t_unit = "s"
|
|
548
|
+
|
|
549
|
+
y_values = t_plot * t_scale
|
|
550
|
+
ylabel = f"Time [{t_unit}]"
|
|
551
|
+
|
|
552
|
+
# Plot spectrogram with frequency on x-axis and time on y-axis
|
|
553
|
+
# Sxx_plot has shape (len(f_plot), len(t_plot)).
|
|
554
|
+
# Transposing Sxx_plot to (len(t_plot), len(f_plot)) matches y-axis (time) and x-axis (frequency).
|
|
555
|
+
mesh = ax.pcolormesh(
|
|
556
|
+
x_values, y_values, Sxx_plot.T, cmap=cmap, shading="auto", **kwargs
|
|
557
|
+
)
|
|
558
|
+
ax.set_xlabel(xlabel)
|
|
559
|
+
ax.set_ylabel(ylabel)
|
|
560
|
+
if title is not None:
|
|
561
|
+
ax.set_title(title)
|
|
562
|
+
|
|
563
|
+
# Add colorbar
|
|
564
|
+
cbar = fig.colorbar(mesh, ax=ax)
|
|
565
|
+
if mode == "psd":
|
|
566
|
+
cbar.set_label("PSD [dB/Hz]")
|
|
567
|
+
elif mode in ("complex", "magnitude"):
|
|
568
|
+
cbar.set_label("Magnitude [dB]")
|
|
569
|
+
elif mode in ("angle", "phase"):
|
|
570
|
+
cbar.set_label("Phase [rad]")
|
|
571
|
+
|
|
572
|
+
if show:
|
|
573
|
+
plt.show()
|
|
574
|
+
return None
|
|
575
|
+
return fig, ax
|