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.
Files changed (84) hide show
  1. commkit/__init__.py +74 -0
  2. commkit/_cuda/__init__.py +321 -0
  3. commkit/_cuda/compiler.py +88 -0
  4. commkit/_cuda/src/bps_min_d2.cu +104 -0
  5. commkit/_cuda/src/cs_block.cu +119 -0
  6. commkit/_cuda/src/selftest.cu +14 -0
  7. commkit/analysis/__init__.py +55 -0
  8. commkit/analysis/_common.py +236 -0
  9. commkit/analysis/allan.py +108 -0
  10. commkit/analysis/drift.py +213 -0
  11. commkit/analysis/interferometry.py +887 -0
  12. commkit/analysis/linewidth.py +480 -0
  13. commkit/analysis/trajectory.py +91 -0
  14. commkit/backend.py +507 -0
  15. commkit/coding/__init__.py +23 -0
  16. commkit/coding/base.py +17 -0
  17. commkit/coding/bch.py +6 -0
  18. commkit/coding/convolutional.py +7 -0
  19. commkit/coding/crc.py +7 -0
  20. commkit/coding/galois.py +8 -0
  21. commkit/coding/hamming.py +6 -0
  22. commkit/coding/interleaving.py +7 -0
  23. commkit/coding/ldpc.py +8 -0
  24. commkit/coding/polar.py +8 -0
  25. commkit/coding/ratematch.py +6 -0
  26. commkit/coding/reed_solomon.py +6 -0
  27. commkit/coding/turbo.py +8 -0
  28. commkit/core/__init__.py +32 -0
  29. commkit/core/frame.py +992 -0
  30. commkit/core/generation.py +581 -0
  31. commkit/core/signal.py +725 -0
  32. commkit/equalization/__init__.py +49 -0
  33. commkit/equalization/_block.py +1855 -0
  34. commkit/equalization/_common.py +606 -0
  35. commkit/equalization/_kernels_jax.py +1720 -0
  36. commkit/equalization/_kernels_numba.py +1704 -0
  37. commkit/equalization/blind.py +223 -0
  38. commkit/equalization/linear.py +365 -0
  39. commkit/equalization/polarization.py +790 -0
  40. commkit/equalization/result.py +191 -0
  41. commkit/equalization/sequential.py +2805 -0
  42. commkit/filtering.py +1120 -0
  43. commkit/frequency.py +1191 -0
  44. commkit/helpers.py +489 -0
  45. commkit/impairments/__init__.py +43 -0
  46. commkit/impairments/channel/__init__.py +20 -0
  47. commkit/impairments/channel/linear.py +310 -0
  48. commkit/impairments/channel/nonlinear.py +11 -0
  49. commkit/impairments/frontend.py +229 -0
  50. commkit/impairments/noise.py +105 -0
  51. commkit/impairments/source.py +219 -0
  52. commkit/io.py +308 -0
  53. commkit/logger.py +103 -0
  54. commkit/mapping/__init__.py +46 -0
  55. commkit/mapping/bits.py +240 -0
  56. commkit/mapping/constellation.py +153 -0
  57. commkit/mapping/gray.py +429 -0
  58. commkit/mapping/llr.py +253 -0
  59. commkit/mapping/shaping.py +218 -0
  60. commkit/metrics.py +949 -0
  61. commkit/multirate.py +476 -0
  62. commkit/plotting/__init__.py +78 -0
  63. commkit/plotting/analysis.py +627 -0
  64. commkit/plotting/constellation.py +483 -0
  65. commkit/plotting/equalizer.py +390 -0
  66. commkit/plotting/eye.py +388 -0
  67. commkit/plotting/spectral.py +575 -0
  68. commkit/plotting/sync.py +953 -0
  69. commkit/plotting/theme.py +203 -0
  70. commkit/plotting/waveform.py +200 -0
  71. commkit/py.typed +0 -0
  72. commkit/recovery/__init__.py +51 -0
  73. commkit/recovery/bps.py +337 -0
  74. commkit/recovery/corrections.py +751 -0
  75. commkit/recovery/pilots.py +803 -0
  76. commkit/recovery/pll.py +482 -0
  77. commkit/recovery/tikhonov.py +424 -0
  78. commkit/recovery/viterbi_viterbi.py +227 -0
  79. commkit/spectral.py +560 -0
  80. commkit/timing.py +841 -0
  81. commkit-1.0.0.dist-info/METADATA +145 -0
  82. commkit-1.0.0.dist-info/RECORD +84 -0
  83. commkit-1.0.0.dist-info/WHEEL +4 -0
  84. 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