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,953 @@
1
+ """Synchronization plots (timing, frequency offset, carrier phase)."""
2
+
3
+ from collections.abc import Sequence
4
+ from typing import Any
5
+
6
+ import matplotlib.pyplot as plt
7
+ import numpy as np
8
+
9
+ from ..backend import to_device
10
+ from .theme import (
11
+ _as_channels,
12
+ _decimate_minmax,
13
+ _grid_figsize,
14
+ _set_eng_formatter,
15
+ )
16
+
17
+
18
+ def plot_timing_correlation(
19
+ corr_mag,
20
+ peak_indices,
21
+ norm_factors,
22
+ threshold: float,
23
+ offset: int = 0,
24
+ ax=None,
25
+ show: bool = False,
26
+ title: str = "Timing Correlation",
27
+ ) -> tuple[Any, Any] | None:
28
+ """
29
+ Plots cross-correlation magnitude for timing estimation diagnostics.
30
+
31
+ For each channel, two panels are drawn: an overall view of the
32
+ correlation and a zoomed view around the detected peak.
33
+
34
+ Parameters
35
+ ----------
36
+ corr_mag : array_like
37
+ Correlation magnitude. Shape: ``(C, N)`` or ``(N,)``.
38
+ peak_indices : array_like
39
+ Integer peak positions per channel. Shape: ``(C,)`` or scalar.
40
+ norm_factors : array_like
41
+ Per-channel normalization factors. The displayed threshold line is
42
+ ``threshold * norm_factors[c]``. Shape: ``(C,)``.
43
+ threshold : float
44
+ Detection threshold (normalized 0-1).
45
+ offset : int, default 0
46
+ Search-range start sample added to sample indices for correct labels.
47
+ ax : array_like of Axes, optional
48
+ Pre-existing axes of shape ``(C, 2)`` - overall and zoom per channel.
49
+ If ``None``, a new figure is created.
50
+ show : bool, default False
51
+ If ``True``, calls ``plt.show()`` and returns ``None``.
52
+ title : str, default "Timing Correlation"
53
+
54
+ Returns
55
+ -------
56
+ (fig, axes) or None
57
+ """
58
+ corr_mag = to_device(corr_mag, "cpu")
59
+ peak_indices = to_device(peak_indices, "cpu").flatten()
60
+ norm_factors = to_device(norm_factors, "cpu").flatten()
61
+
62
+ if corr_mag.ndim == 1:
63
+ corr_mag = corr_mag[None, :]
64
+ C = corr_mag.shape[0]
65
+
66
+ if ax is None:
67
+ fig, axes = plt.subplots(C, 2, figsize=_grid_figsize(C, 2), squeeze=False)
68
+ else:
69
+ axes = ax
70
+ fig = axes[0][0].figure
71
+
72
+ for i in range(C):
73
+ ax1 = axes[i][0]
74
+ ax2 = axes[i][1]
75
+
76
+ c_ch = corr_mag[i]
77
+ pk_idx = int(peak_indices[i]) if i < len(peak_indices) else 0
78
+ norm_val = float(norm_factors[i]) if i < len(norm_factors) else 1.0
79
+ abs_thresh = threshold * norm_val
80
+ metric_val = float(c_ch[pk_idx] / norm_val) if norm_val > 0 else 0.0
81
+ ch_suffix = f" - Ch {i}" if C > 1 else ""
82
+
83
+ x_all = np.arange(len(c_ch)) + offset
84
+ ax1.plot(x_all, c_ch, label=f"Ch {i} metric={metric_val:.2f}")
85
+ ax1.axhline(float(c_ch[pk_idx]), color="r", linestyle="--", label="Peak")
86
+ ax1.axvline(pk_idx + offset, color="r", linestyle="--")
87
+ if abs_thresh > 0:
88
+ ax1.axhline(abs_thresh, color="g", linestyle=":", label="Thresh")
89
+ ax1.set_title(f"{title}{ch_suffix}")
90
+ ax1.set_xlabel("Sample Index")
91
+ ax1.set_ylabel("|R|")
92
+ ax1.legend(loc="upper right")
93
+
94
+ zoom_w = 40
95
+ s_z = max(0, pk_idx - zoom_w)
96
+ e_z = min(len(c_ch), pk_idx + zoom_w)
97
+ x_zoom = np.arange(s_z, e_z) + offset
98
+ ax2.plot(x_zoom, c_ch[s_z:e_z], label="Peak area")
99
+ ax2.axvline(
100
+ pk_idx + offset, color="r", linestyle="--", label=f"Pk @ {pk_idx + offset}"
101
+ )
102
+ if abs_thresh > 0:
103
+ ax2.axhline(abs_thresh, color="g", linestyle=":", label="Thresh")
104
+ ax2.set_title(f"{title}{ch_suffix} - Detail")
105
+ ax2.set_xlabel("Sample Index")
106
+ ax2.set_ylabel("|R|")
107
+ ax2.legend(loc="upper right")
108
+
109
+ if show:
110
+ plt.show()
111
+ return None
112
+ return fig, axes
113
+
114
+
115
+ def plot_mm_autocorrelation(
116
+ R_np,
117
+ f_est,
118
+ sampling_rate: float,
119
+ M: int = 1,
120
+ ax=None,
121
+ show: bool = False,
122
+ title: str = "FOE - Mengali-Morelli",
123
+ ) -> tuple[Any, Any] | None:
124
+ """
125
+ Plots the Mengali-Morelli autocorrelation diagnostics.
126
+
127
+ Per-channel two-panel layout:
128
+
129
+ * **Left / Top** - Normalised autocorrelation magnitude ``|R[m]|`` vs lag
130
+ ``m``. Encodes per-lag SNR; used as weight proxy ``w[m] ∝ m²|R[m]|²``.
131
+ * **Right / Bottom** - Wrapped phase ``angle(R[m])`` vs lag, with the
132
+ expected linear ramp overlaid and ``±π`` wrap boundaries marked.
133
+
134
+ Parameters
135
+ ----------
136
+ R_np : (L,) or (C, L) complex128
137
+ Normalised per-channel autocorrelation at lags ``m = 1 ... L``.
138
+ A 1-D input is treated as a single channel.
139
+ f_est : float or list of float
140
+ Frequency offset estimate(s) in Hz. Scalar for a single channel;
141
+ list/array of length ``C`` for multi-channel input.
142
+ sampling_rate : float
143
+ Sampling rate in Hz.
144
+ M : int, default 1
145
+ Modulation pre-processing exponent.
146
+ ax : Axes or array of Axes, optional
147
+ For a single channel: pair ``[ax_amp, ax_phase]``.
148
+ For multiple channels: array of shape ``(C, 2)``.
149
+ A new figure is created when ``None``.
150
+ show : bool, default False
151
+ If ``True``, calls ``plt.show()`` and returns ``None``.
152
+ title : str, default "FOE - Mengali-Morelli"
153
+
154
+ Returns
155
+ -------
156
+ (fig, axes) or None
157
+ Single channel: ``axes`` is ``[ax_amp, ax_phase]``.
158
+ Multi-channel: ``axes`` is a list of ``[ax_amp_c, ax_phase_c]`` pairs.
159
+ """
160
+ R_np = np.asarray(R_np)
161
+ if R_np.ndim == 1:
162
+ R_np = R_np[None, :] # (1, L)
163
+ C, L = R_np.shape
164
+
165
+ if hasattr(f_est, "__len__"):
166
+ f_ests = [float(f) for f in f_est]
167
+ else:
168
+ f_ests = [float(f_est)] * C
169
+
170
+ lags = np.arange(1, L + 1)
171
+
172
+ if ax is None:
173
+ if C == 1:
174
+ fig, raw_axes = plt.subplots(1, 2, figsize=_grid_figsize(1, 2))
175
+ axes_per_ch = [raw_axes] # list of (ax_amp, ax_phase)
176
+ else:
177
+ fig, raw_axes = plt.subplots(
178
+ C, 2, figsize=_grid_figsize(C, 2), squeeze=False
179
+ )
180
+ axes_per_ch = [(raw_axes[c, 0], raw_axes[c, 1]) for c in range(C)]
181
+ else:
182
+ ax_arr = np.asarray(ax, dtype=object)
183
+ if ax_arr.ndim == 1 and ax_arr.shape[0] == 2:
184
+ axes_per_ch = [ax_arr]
185
+ else:
186
+ axes_per_ch = [ax_arr[c] for c in range(C)]
187
+ fig = axes_per_ch[0][0].figure
188
+
189
+ for c in range(C):
190
+ ax_amp, ax_phase = axes_per_ch[c]
191
+ amp = np.abs(R_np[c])
192
+ theta = np.angle(R_np[c])
193
+ f_c = f_ests[c]
194
+ ch_suffix = f" - Ch {c}" if C > 1 else ""
195
+ expected_phase = 2.0 * np.pi * f_c * M * lags / sampling_rate
196
+
197
+ ax_amp.plot(lags, amp, color=f"C{c}")
198
+ ax_amp.set_xlabel("Lag m")
199
+ ax_amp.set_ylabel("|R[m]|")
200
+ ax_amp.set_title(f"{title}{ch_suffix} ($\\Delta f$={f_c:.2f} Hz, M={M})")
201
+
202
+ ax_phase.scatter(
203
+ lags, theta, s=6, color=f"C{c}", label="Angle(R[m]) (wrapped)", zorder=3
204
+ )
205
+ ax_phase.plot(
206
+ lags,
207
+ (expected_phase + np.pi) % (2 * np.pi) - np.pi,
208
+ color="red",
209
+ linestyle="--",
210
+ label=f"Expected $2\\pi \\cdot \\Delta f \\cdot M \\cdot m/f_s$ "
211
+ f"($\\Delta f$={f_c:.2f} Hz)",
212
+ )
213
+ ax_phase.axhline(
214
+ np.pi, color="gray", linestyle=":", label=r"$\pm\pi$ wrap boundary"
215
+ )
216
+ ax_phase.axhline(-np.pi, color="gray", linestyle=":")
217
+ ax_phase.set_xlabel("Lag m")
218
+ ax_phase.set_ylabel("Phase [rad]")
219
+ ax_phase.set_ylim(-np.pi - 0.3, np.pi + 0.3)
220
+ ax_phase.legend()
221
+
222
+ if show:
223
+ plt.show()
224
+ return None
225
+ return fig, (axes_per_ch[0] if C == 1 else axes_per_ch)
226
+
227
+
228
+ def plot_frequency_offset_spectrum(
229
+ mag_spectrum,
230
+ freqs,
231
+ M: int,
232
+ k_peaks,
233
+ f_estimates,
234
+ search_range=None,
235
+ ax=None,
236
+ show: bool = False,
237
+ title: str = "FOE - M-th Power Spectrum",
238
+ ) -> tuple[Any, Any] | None:
239
+ """
240
+ Plots the M-th power spectrum used for blind frequency offset estimation.
241
+
242
+ The spectrum has a tone at ``M·Δf``; this function maps the x-axis back
243
+ to ``Δf`` by dividing by ``M`` so the detected peak aligns with the
244
+ reported frequency estimate.
245
+
246
+ Parameters
247
+ ----------
248
+ mag_spectrum : array_like
249
+ Magnitude spectrum ``|X^M(f)|``. Shape: ``(C, nfft)`` or ``(nfft,)``.
250
+ freqs : array_like
251
+ Frequency axis in Hz (``np.fft.fftfreq(nfft) * fs``).
252
+ Shape: ``(nfft,)``.
253
+ M : int
254
+ M-th power used to remove modulation.
255
+ k_peaks : array_like
256
+ Peak bin index per channel. Shape: ``(C,)`` or scalar.
257
+ f_estimates : list of float
258
+ Per-channel frequency estimates in Hz (before multi-channel averaging).
259
+ search_range : tuple of float, optional
260
+ ``(f_min, f_max)`` in Hz that restricted the search. Shown as a
261
+ shaded region.
262
+ ax : array_like of Axes, optional
263
+ One Axes per channel. If ``None``, a new figure is created.
264
+ show : bool, default False
265
+ title : str, default "FOE - M-th Power Spectrum"
266
+
267
+ Returns
268
+ -------
269
+ (fig, axes) or None
270
+ """
271
+ mag_spectrum = to_device(mag_spectrum, "cpu")
272
+ freqs = to_device(freqs, "cpu")
273
+ k_peaks = to_device(k_peaks, "cpu").flatten()
274
+ f_estimates = list(f_estimates)
275
+
276
+ if mag_spectrum.ndim == 1:
277
+ mag_spectrum = mag_spectrum[None, :]
278
+ C = mag_spectrum.shape[0]
279
+
280
+ if ax is None:
281
+ fig, raw_axes = plt.subplots(1, C, figsize=_grid_figsize(1, C), squeeze=False)
282
+ axes_list = list(raw_axes[0])
283
+ else:
284
+ axes_list = list(ax) if hasattr(ax, "__len__") else [ax]
285
+ fig = axes_list[0].figure
286
+
287
+ sort_idx = np.argsort(freqs)
288
+ f_sorted = freqs[sort_idx]
289
+ f_delta = f_sorted / M # map M·Δf -> Δf
290
+
291
+ for i in range(C):
292
+ axi = axes_list[i]
293
+ mag = mag_spectrum[i][sort_idx]
294
+ f_est = float(f_estimates[i]) if i < len(f_estimates) else 0.0
295
+ ch_suffix = f" - Ch {i}" if C > 1 else ""
296
+
297
+ axi.plot(f_delta, mag, color="C0")
298
+ axi.axvline(
299
+ f_est,
300
+ color="r",
301
+ linestyle="--",
302
+ label=rf"$\hat{{f}}$ = {f_est:.2f} Hz",
303
+ )
304
+ if search_range is not None:
305
+ axi.axvspan(
306
+ search_range[0],
307
+ search_range[1],
308
+ alpha=0.4,
309
+ color="green",
310
+ label="Search range",
311
+ )
312
+ axi.set_title(f"{title}{ch_suffix}")
313
+ axi.set_xlabel(f"$\\Delta f$ [Hz] (÷M={M} applied)")
314
+ axi.set_ylabel(f"|X^{M}(f)|")
315
+ axi.legend()
316
+
317
+ if show:
318
+ plt.show()
319
+ return None
320
+ return fig, (axes_list[0] if C == 1 else axes_list)
321
+
322
+
323
+ def plot_carrier_phase_trajectory(
324
+ phi_full,
325
+ block_centers=None,
326
+ phi_blocks=None,
327
+ n_train: int = 0,
328
+ ax=None,
329
+ show: bool = False,
330
+ title: str = "Carrier Phase Trajectory",
331
+ ) -> tuple[Any, Any] | None:
332
+ """
333
+ Plots per-symbol carrier phase trajectory for CPR algorithm diagnostics.
334
+
335
+ All channels are overlaid on a single subplot. Block-based methods may
336
+ pass ``block_centers`` / ``phi_blocks`` to annotate block-phase estimates
337
+ as vertical lines (not scatter markers, consistent with the joint
338
+ equalizer phase panel).
339
+
340
+ Parameters
341
+ ----------
342
+ phi_full : array_like
343
+ Per-symbol phase estimate in radians. Shape: ``(C, N)`` or ``(N,)``.
344
+ block_centers : array_like, optional
345
+ Block centre positions in symbols (VV, BPS). Shape: ``(N_blocks,)``.
346
+ If provided, thin vertical lines at each block centre are drawn.
347
+ phi_blocks : array_like, optional
348
+ Kept for backwards compatibility - ignored (block markers removed).
349
+ n_train : int, default 0
350
+ Training/DD boundary symbol index. Draws a dashed vertical line.
351
+ ax : Axes, optional
352
+ Single Axes object to plot into. If ``None``, a new figure is created.
353
+ show : bool, default False
354
+ title : str, default "Carrier Phase Trajectory"
355
+
356
+ Returns
357
+ -------
358
+ (fig, ax) or None
359
+ """
360
+ phi_full = to_device(phi_full, "cpu")
361
+ if phi_full.ndim == 1:
362
+ phi_full = phi_full[None, :]
363
+ C, N = phi_full.shape
364
+
365
+ if ax is None:
366
+ fig, axi = plt.subplots(1, 1)
367
+ else:
368
+ axi = ax
369
+ fig = axi.figure
370
+
371
+ sym_idx = np.arange(N)
372
+ for i in range(C):
373
+ phi_deg = np.degrees(phi_full[i])
374
+ label = f"Ch {i}" if C > 1 else None
375
+ axi.plot(sym_idx, phi_deg, label=label)
376
+
377
+ if block_centers is not None:
378
+ bc = np.asarray(to_device(block_centers, "cpu"))
379
+ for tc in bc:
380
+ axi.axvline(tc, color="gray", alpha=0.4)
381
+
382
+ if n_train > 0:
383
+ axi.axvline(
384
+ n_train,
385
+ color="white",
386
+ linestyle="--",
387
+ label=f"DD start ({n_train})",
388
+ )
389
+
390
+ phi_all = np.degrees(phi_full)
391
+ phi_mean = float(np.mean(phi_all))
392
+ phi_std = float(np.std(phi_all))
393
+ axi.set_title(f"{title} [$\\mu$={phi_mean:.1f}°, $\\sigma$={phi_std:.2f}°]")
394
+ axi.set_xlabel("Symbol Index")
395
+ axi.set_ylabel("Phase [deg]")
396
+ if C > 1 or n_train > 0:
397
+ axi.legend(loc="upper right")
398
+
399
+ if show:
400
+ plt.show()
401
+ return None
402
+ return fig, axi
403
+
404
+
405
+ def plot_frequency_offset_blockwise_result(
406
+ t_centers,
407
+ df_estimates,
408
+ n_grid,
409
+ df_dense,
410
+ phase_trajectory,
411
+ ax=None,
412
+ show: bool = False,
413
+ title: str = "Block-wise FOE",
414
+ max_points: int = 4000,
415
+ ) -> tuple[Any, Any] | None:
416
+ """
417
+ Diagnostic plot for ``frequency.correct_frequency_offset_blockwise``.
418
+
419
+ Shows three panels:
420
+
421
+ 1. Per-block frequency estimates (scatter) and interpolated Δf trajectory (line).
422
+ 2. Integrated phase trajectory in degrees.
423
+
424
+ Parameters
425
+ ----------
426
+ t_centers : array_like
427
+ Block centre sample indices. Shape: ``(K,)``.
428
+ df_estimates : array_like
429
+ Per-block frequency estimates in Hz. Shape: ``(K,)``.
430
+ n_grid : array_like
431
+ Dense sample index grid. Shape: ``(N,)``.
432
+ df_dense : array_like
433
+ Interpolated frequency at each sample in Hz. Shape: ``(N,)``.
434
+ phase_trajectory : array_like
435
+ Integrated phase in radians. Shape: ``(N,)``.
436
+ ax : list of 2 Axes, optional
437
+ Pre-existing axes ``[ax_freq, ax_phase]``. If ``None``, a new figure
438
+ with 2 panels is created.
439
+ show : bool, default False
440
+ title : str
441
+ max_points : int, default 4000
442
+ Approximate per-trace point budget; the dense per-sample Δf and phase
443
+ traces are envelope-decimated (min/max) before plotting so long records
444
+ render quickly. Pass ``<= 0`` to plot every point.
445
+
446
+ Returns
447
+ -------
448
+ (fig, axes) or None
449
+ """
450
+ t_centers = np.asarray(t_centers, dtype=np.float64)
451
+ df_estimates = np.asarray(df_estimates, dtype=np.float64)
452
+ n_grid = np.asarray(n_grid, dtype=np.float64)
453
+ df_dense = np.asarray(df_dense, dtype=np.float64)
454
+ phase_trajectory = np.asarray(phase_trajectory, dtype=np.float64)
455
+
456
+ if ax is None:
457
+ fig, axes = plt.subplots(1, 2, figsize=_grid_figsize(1, 2))
458
+ else:
459
+ axes = np.asarray(ax).flatten()[:2]
460
+ fig = axes[0].figure
461
+
462
+ # Decimate the dense per-sample traces for fast rendering (envelope-preserving).
463
+ nf, df_k = _decimate_minmax(n_grid, df_dense * 1e-3, max_points)
464
+ np_, ph_deg = _decimate_minmax(n_grid, np.degrees(phase_trajectory), max_points)
465
+
466
+ # Panel 1: frequency trajectory
467
+ ax_f = axes[0]
468
+ ax_f.plot(nf, df_k, label=r"Interpolated $\Delta f$")
469
+ ax_f.scatter(
470
+ t_centers,
471
+ df_estimates * 1e-3,
472
+ s=50,
473
+ zorder=5,
474
+ color="C1",
475
+ label="Block estimate",
476
+ )
477
+ ax_f.set_xlabel("Sample Index")
478
+ ax_f.set_ylabel(r"$\Delta f$ [kHz]")
479
+ ax_f.set_title(f"{title} - Frequency")
480
+ ax_f.legend()
481
+
482
+ # Panel 2: phase trajectory
483
+ ax_p = axes[1]
484
+ ax_p.plot(np_, ph_deg)
485
+ ax_p.set_xlabel("Sample Index")
486
+ ax_p.set_ylabel("Phase [deg]")
487
+ ax_p.set_title(f"{title} - Integrated Phase")
488
+
489
+ if show:
490
+ plt.show()
491
+ return None
492
+ return fig, axes
493
+
494
+
495
+ def plot_pilot_phase_estimate(
496
+ pilot_indices,
497
+ phi_pilots_u,
498
+ phi_full=None,
499
+ f_est: float | Sequence[float] | np.ndarray = 0.0,
500
+ sampling_rate: float = 1.0,
501
+ ax=None,
502
+ show: bool = False,
503
+ title: str = "Pilot Phase Estimate",
504
+ ) -> tuple[Any, Any] | None:
505
+ """
506
+ Plots pilot phase scatter, linear fit, and the full interpolated trajectory.
507
+
508
+ Used as a diagnostic for both pilot-based frequency offset estimation
509
+ (``estimate_frequency_offset_pilot_symbols``) and pilot-aided carrier phase
510
+ recovery (``recover_carrier_phase_pilot_symbols``).
511
+
512
+ Parameters
513
+ ----------
514
+ pilot_indices : array_like
515
+ Sample indices of pilot positions. Shape: ``(P,)``.
516
+ phi_pilots_u : array_like
517
+ Unwrapped pilot phases in radians. Shape: ``(C, P)`` or ``(P,)``.
518
+ phi_full : array_like, optional
519
+ Interpolated per-symbol phase in radians. Shape: ``(C, N)`` or ``(N,)``.
520
+ If provided, a second panel shows the full trajectory per channel.
521
+ f_est : float or list of float, default 0.0
522
+ Estimated frequency offset in Hz (annotation on the fit line).
523
+ Scalar applies the same label to all channels; list of length ``C``
524
+ annotates each channel independently.
525
+ sampling_rate : float, default 1.0
526
+ Sampling rate in Hz.
527
+ ax : array_like of Axes, optional
528
+ Shape ``(C, 1)`` when ``phi_full`` is ``None``, else ``(C, 2)``.
529
+ If ``None``, a new figure is created.
530
+ show : bool, default False
531
+ title : str, default "Pilot Phase Estimate"
532
+
533
+ Returns
534
+ -------
535
+ (fig, axes) or None
536
+ """
537
+ pilot_indices = to_device(pilot_indices, "cpu").astype(float)
538
+ phi_pilots_u = to_device(phi_pilots_u, "cpu")
539
+ if phi_pilots_u.ndim == 1:
540
+ phi_pilots_u = phi_pilots_u[None, :]
541
+ C, P = phi_pilots_u.shape
542
+
543
+ if isinstance(f_est, (np.ndarray, list, tuple)):
544
+ f_ests = [float(f) for f in f_est] # type: ignore[union-attr]
545
+ else:
546
+ f_ests = [float(f_est)] * C # type: ignore[arg-type]
547
+
548
+ has_full = phi_full is not None
549
+ if has_full:
550
+ phi_full = to_device(phi_full, "cpu")
551
+ if phi_full.ndim == 1:
552
+ phi_full = phi_full[None, :]
553
+ N = phi_full.shape[1]
554
+
555
+ n_cols = 2 if has_full else 1
556
+ axes: Any = None
557
+ if ax is None:
558
+ fig, axes = plt.subplots(
559
+ C, n_cols, figsize=_grid_figsize(C, n_cols), squeeze=False
560
+ )
561
+ else:
562
+ # ax expected as (C, n_cols) sequence - use list of lists to avoid
563
+ # np.reshape() which fails when ax contains non-array objects
564
+ if hasattr(ax, "__len__") and hasattr(ax[0], "__len__"):
565
+ axes = ax # already 2-D list
566
+ elif hasattr(ax, "__len__"):
567
+ axes = [[a] for a in ax] # 1-D list -> wrap each in a row
568
+ else:
569
+ axes = [[ax]]
570
+ fig = axes[0][0].figure
571
+
572
+ t_pilots = pilot_indices / sampling_rate
573
+ for i in range(C):
574
+ phi_p = phi_pilots_u[i]
575
+ ch_suffix = f" - Ch {i}" if C > 1 else ""
576
+
577
+ # Linear fit: φ(t) = 2π·Δf·t + φ₀
578
+ if P > 1:
579
+ t_c = t_pilots - np.mean(t_pilots)
580
+ t_var = float(np.dot(t_c, t_c))
581
+ slope = (
582
+ float(np.dot(t_c, phi_p - np.mean(phi_p)) / t_var) if t_var > 0 else 0.0
583
+ )
584
+ phi_fit = slope * t_pilots + (np.mean(phi_p) - slope * np.mean(t_pilots))
585
+ else:
586
+ phi_fit = phi_p.copy()
587
+
588
+ ax1 = axes[i][0]
589
+ ax1.scatter(
590
+ pilot_indices,
591
+ np.degrees(phi_p),
592
+ s=14,
593
+ zorder=5,
594
+ label=r"Pilot $\phi$ (unwrapped)",
595
+ )
596
+ ax1.plot(
597
+ pilot_indices,
598
+ np.degrees(phi_fit),
599
+ "r--",
600
+ label=f"Fit $\\Delta f$={f_ests[i]:.3f} Hz",
601
+ )
602
+ ax1.set_title(f"{title}{ch_suffix} - Pilots")
603
+ ax1.set_xlabel("Sample Index")
604
+ ax1.set_ylabel("Phase [deg]")
605
+ ax1.legend()
606
+
607
+ if has_full:
608
+ ax2 = axes[i][1]
609
+ sym_idx = np.arange(N)
610
+ ax2.plot(
611
+ sym_idx,
612
+ np.degrees(phi_full[i]),
613
+ alpha=0.4,
614
+ label=r"$\hat{\phi}$ (interpolated)",
615
+ )
616
+ ax2.scatter(
617
+ pilot_indices,
618
+ np.degrees(phi_p),
619
+ s=10,
620
+ color="r",
621
+ zorder=5,
622
+ label="Pilots",
623
+ )
624
+ ax2.set_title(f"{title}{ch_suffix} - Full Trajectory")
625
+ ax2.set_xlabel("Symbol Index")
626
+ ax2.set_ylabel("Phase [deg]")
627
+ ax2.legend()
628
+
629
+ if show:
630
+ plt.show()
631
+ return None
632
+ return fig, axes
633
+
634
+
635
+ def plot_pilot_tone_phase_estimate(
636
+ freqs,
637
+ mag_spectrum,
638
+ window,
639
+ f_tones,
640
+ theta,
641
+ tone_frequency: float,
642
+ bandwidth: float,
643
+ ax=None,
644
+ show: bool = False,
645
+ title: str = "CPR - Pilot Tone",
646
+ max_points: int = 4000,
647
+ ) -> tuple[Any, Any] | None:
648
+ """
649
+ Diagnostic for ``recover_carrier_phase_pilot_tone``.
650
+
651
+ Two panels:
652
+
653
+ 1. **Tone spectrum** - magnitude spectrum ``|X(f)|`` (dB) with the
654
+ zero-phase extraction window overlaid on a twin axis, the per-channel
655
+ refined tone peak marked, and the nominal tone frequency annotated.
656
+ 2. **Recovered phase** - the per-sample phase estimate ``θ̂[n]`` (deg).
657
+
658
+ All channels are overlaid on each panel.
659
+
660
+ Parameters
661
+ ----------
662
+ freqs : array_like
663
+ FFT frequency axis in Hz (``np.fft.fftfreq(N) * fs``). Shape: ``(N,)``.
664
+ mag_spectrum : array_like
665
+ Magnitude spectrum ``|X(f)|``. Shape: ``(C, N)`` or ``(N,)``.
666
+ window : array_like
667
+ Extraction window ``W(f)`` in ``[0, 1]``. Shape: ``(C, N)`` or ``(N,)``.
668
+ f_tones : array_like
669
+ Per-channel refined tone frequency in Hz. Shape: ``(C,)``.
670
+ theta : array_like
671
+ Recovered per-sample phase in radians. Shape: ``(C, N)`` or ``(N,)``.
672
+ tone_frequency : float
673
+ Nominal pilot-tone frequency in Hz.
674
+ bandwidth : float
675
+ Extraction window half-width in Hz (shaded around the nominal tone).
676
+ ax : array_like of Axes, optional
677
+ Two Axes ``[spectrum, phase]``. If ``None``, a new figure is created.
678
+ show : bool, default False
679
+ title : str, default "CPR - Pilot Tone"
680
+ max_points : int, default 4000
681
+ Approximate per-trace point budget; longer traces are envelope-decimated
682
+ (min/max) before plotting so oversampled records render quickly without
683
+ losing the tone peak or window edges. Pass ``<= 0`` to plot every point.
684
+
685
+ Returns
686
+ -------
687
+ (fig, axes) or None
688
+ """
689
+ freqs = np.asarray(to_device(freqs, "cpu"), dtype=float)
690
+ mag_spectrum = to_device(mag_spectrum, "cpu")
691
+ window = to_device(window, "cpu")
692
+ theta = to_device(theta, "cpu")
693
+ if mag_spectrum.ndim == 1:
694
+ mag_spectrum = mag_spectrum[None, :]
695
+ if window.ndim == 1:
696
+ window = window[None, :]
697
+ if theta.ndim == 1:
698
+ theta = theta[None, :]
699
+ f_tones = np.atleast_1d(np.asarray(to_device(f_tones, "cpu"), dtype=float))
700
+ C, N = mag_spectrum.shape
701
+
702
+ if ax is None:
703
+ fig, raw_axes = plt.subplots(1, 2, figsize=_grid_figsize(1, 2), squeeze=False)
704
+ ax_spec, ax_phase = raw_axes[0]
705
+ else:
706
+ ax_spec, ax_phase = ax[0], ax[1]
707
+ fig = ax_spec.figure
708
+
709
+ order = np.argsort(freqs)
710
+ f_sorted = freqs[order]
711
+ eps = 1e-300
712
+
713
+ # Panel 1 - spectrum (dB) with extraction window on a twin axis.
714
+ # Envelope-decimate the full-resolution traces: min/max preserves the tone
715
+ # peak and the window edges that plain striding would skip over.
716
+ ax_win = ax_spec.twinx()
717
+ for i in range(C):
718
+ mag_db = 20.0 * np.log10(np.maximum(mag_spectrum[i][order], eps))
719
+ label = f"Ch {i}" if C > 1 else "|X(f)|"
720
+ f_d, mag_d = _decimate_minmax(f_sorted, mag_db, max_points)
721
+ f_w, win_d = _decimate_minmax(f_sorted, window[i][order], max_points)
722
+ ax_spec.plot(f_d, mag_d, label=label)
723
+ ax_win.plot(f_w, win_d, color="C3", alpha=0.4)
724
+ ax_spec.axvline(
725
+ float(f_tones[i]),
726
+ color="C2",
727
+ linestyle=":",
728
+ label=("Tone peak" if i == 0 else None),
729
+ )
730
+ ax_spec.axvspan(
731
+ tone_frequency - bandwidth,
732
+ tone_frequency + bandwidth,
733
+ color="C3",
734
+ alpha=0.4,
735
+ label=f"Window ±B ({bandwidth:.3g} Hz)",
736
+ )
737
+ ax_spec.axvline(
738
+ tone_frequency,
739
+ color="k",
740
+ linestyle="--",
741
+ label=f"Nominal f_p ({tone_frequency:.3g} Hz)",
742
+ )
743
+ ax_win.set_ylim(-0.05, 1.35)
744
+ ax_win.set_ylabel("Window W(f)", color="C3")
745
+ ax_spec.set_title(f"{title} - Tone Spectrum")
746
+ ax_spec.set_xlabel("Frequency [Hz]")
747
+ ax_spec.set_ylabel("|X(f)| [dB]")
748
+ ax_spec.legend(loc="upper left")
749
+
750
+ # Panel 2 - recovered phase trajectory.
751
+ sample_idx = np.arange(N)
752
+ for i in range(C):
753
+ ph_label = f"Ch {i}" if C > 1 else None
754
+ n_d, th_d = _decimate_minmax(sample_idx, np.degrees(theta[i]), max_points)
755
+ ax_phase.plot(n_d, th_d, label=ph_label)
756
+ th_mean = float(np.mean(np.degrees(theta)))
757
+ th_std = float(np.std(np.degrees(theta)))
758
+ ax_phase.set_title(
759
+ f"{title} - Recovered $\\hat{{\\theta}}$ "
760
+ f"[$\\mu$={th_mean:.1f}°, $\\sigma$={th_std:.2f}°]"
761
+ )
762
+ ax_phase.set_xlabel("Sample Index")
763
+ ax_phase.set_ylabel("Phase [deg]")
764
+ if C > 1:
765
+ ax_phase.legend(loc="upper right")
766
+
767
+ if show:
768
+ plt.show()
769
+ return None
770
+ return fig, (ax_spec, ax_phase)
771
+
772
+
773
+ def plot_pilot_tones_phase_estimate(
774
+ delta,
775
+ phi,
776
+ ref: int,
777
+ used: Sequence[int],
778
+ ax=None,
779
+ show: bool = False,
780
+ title: str = "CPR - Pilot Tones (MRC)",
781
+ max_points: int = 4000,
782
+ ) -> tuple[Any, Any] | None:
783
+ """
784
+ Diagnostic for ``recover_carrier_phase_pilot_tones``.
785
+
786
+ Two panels:
787
+
788
+ 1. **Inter-tone differential** - the slow tracked phase ``δ_k[n]`` (deg) for
789
+ each non-reference tone; the line style flags whether the tone was
790
+ combined (solid) or gated out (dashed).
791
+ 2. **Combined phase** - the per-sample common estimate ``φ̂[n]`` (deg).
792
+
793
+ All channels share the active theme; no colours or line widths are set
794
+ explicitly, so the trace colours come from the rcParams cycle.
795
+
796
+ Parameters
797
+ ----------
798
+ delta : sequence of array_like
799
+ Per-tone differential phase ``δ_k[n]`` in radians - a length-``K``
800
+ sequence (or ``(K, N)`` array) of ``(N,)`` traces. The reference tone's
801
+ entry is ``≈ 0`` and is skipped in panel 1.
802
+ phi : array_like
803
+ Combined per-sample phase estimate ``φ̂[n]`` in radians. Shape ``(N,)``.
804
+ ref : int
805
+ Reference-tone index (its differential is identically zero).
806
+ used : sequence of int
807
+ Indices of the tones that were combined (the rest were gated out).
808
+ ax : array_like of Axes, optional
809
+ Two Axes ``[differential, phase]``. If ``None``, a new figure is created.
810
+ show : bool, default False
811
+ title : str, default "CPR - Pilot Tones (MRC)"
812
+ max_points : int, default 4000
813
+ Per-trace point budget; longer traces are envelope-decimated (min/max)
814
+ before plotting. Pass ``<= 0`` to plot every point.
815
+
816
+ Returns
817
+ -------
818
+ (fig, axes) or None
819
+ """
820
+ delta = [np.asarray(to_device(d, "cpu"), dtype=float) for d in delta]
821
+ phi = np.asarray(to_device(phi, "cpu"), dtype=float)
822
+ K = len(delta)
823
+ used_set = {int(u) for u in used}
824
+
825
+ if ax is None:
826
+ fig, raw_axes = plt.subplots(1, 2, figsize=_grid_figsize(1, 2), squeeze=False)
827
+ ax_delta, ax_phase = raw_axes[0]
828
+ else:
829
+ ax_delta, ax_phase = ax[0], ax[1]
830
+ fig = ax_delta.figure
831
+
832
+ # Panel 1 - per-tone slow differential (skip the reference, whose δ ≡ 0).
833
+ n_plotted = 0
834
+ for k in range(K):
835
+ if k == ref:
836
+ continue
837
+ gated = k not in used_set
838
+ idx = np.arange(delta[k].shape[-1])
839
+ n_d, d_d = _decimate_minmax(idx, np.degrees(delta[k]), max_points)
840
+ ax_delta.plot(
841
+ n_d,
842
+ d_d,
843
+ linestyle="--" if gated else "-",
844
+ label=f"$\\delta$ tone {k}" + (" (gated)" if gated else ""),
845
+ )
846
+ n_plotted += 1
847
+ ax_delta.set_title(f"{title} - Inter-tone $\\delta$ (ref = tone {ref})")
848
+ ax_delta.set_xlabel("Sample Index")
849
+ ax_delta.set_ylabel("Differential phase [deg]")
850
+ if n_plotted:
851
+ ax_delta.legend(loc="upper right")
852
+
853
+ # Panel 2 - combined common-phase track.
854
+ idx = np.arange(phi.shape[-1])
855
+ n_d, p_d = _decimate_minmax(idx, np.degrees(phi), max_points)
856
+ ax_phase.plot(n_d, p_d)
857
+ ph_std = float(np.std(np.degrees(phi)))
858
+ ax_phase.set_title(
859
+ f"{title} - Combined $\\hat{{\\phi}}$ [$\\sigma$={ph_std:.2f}°, used={list(used)}]"
860
+ )
861
+ ax_phase.set_xlabel("Sample Index")
862
+ ax_phase.set_ylabel("Phase [deg]")
863
+
864
+ if show:
865
+ plt.show()
866
+ return None
867
+ return fig, (ax_delta, ax_phase)
868
+
869
+
870
+ def plot_carrier_phase_decomposition(
871
+ phi,
872
+ drift=None,
873
+ *,
874
+ symbol_rate: float,
875
+ n_train: int = 0,
876
+ ax=None,
877
+ show: bool = False,
878
+ title: str = "Recovered carrier phase",
879
+ ) -> tuple[Any, Any] | None:
880
+ """
881
+ Plots the recovered carrier-phase trajectory and its slow drift component.
882
+
883
+ The total unwrapped phase phi(t) is drawn faintly with the
884
+ low-pass drift overlaid in bold, visualising the
885
+ ``analysis.separate_drift_phase_noise`` split. MIMO inputs overlay all
886
+ channels.
887
+
888
+ Parameters
889
+ ----------
890
+ phi : array_like
891
+ Unwrapped carrier phase in radians. Shape ``(N,)`` or ``(C, N)``.
892
+ drift : array_like, optional
893
+ Drift (low-pass) component, same shape as ``phi``. Overlaid in bold.
894
+ symbol_rate : float
895
+ Symbol rate in Baud; sets the (seconds) time axis.
896
+ n_train : int, default 0
897
+ If > 0, draws a dashed training/DD boundary marker.
898
+ ax : Axes, optional
899
+ Target axes; a new figure is created when None.
900
+ show : bool, default False
901
+ If True, calls ``plt.show()`` and returns None.
902
+ title : str
903
+
904
+ Returns
905
+ -------
906
+ (fig, ax) or None
907
+ """
908
+ phi_c = _as_channels(phi)
909
+ C, N = phi_c.shape
910
+ drift_c = _as_channels(drift) if drift is not None else None
911
+
912
+ if ax is None:
913
+ fig, axi = plt.subplots(1, 1)
914
+ else:
915
+ axi = ax
916
+ fig = axi.figure
917
+
918
+ t = np.arange(N) / float(symbol_rate)
919
+ for i in range(C):
920
+ clabel = f"pol {i}" if C > 1 else r"$\phi$ total"
921
+ axi.plot(
922
+ t,
923
+ phi_c[i],
924
+ color=f"C{i}",
925
+ alpha=0.4,
926
+ label=clabel if drift_c is None else None,
927
+ )
928
+ if drift_c is not None:
929
+ axi.plot(
930
+ t,
931
+ drift_c[i],
932
+ color=f"C{i}",
933
+ label=f"Drift (pol {i})" if C > 1 else "Drift",
934
+ )
935
+
936
+ if n_train > 0:
937
+ axi.axvline(
938
+ n_train / float(symbol_rate),
939
+ color="white",
940
+ ls="--",
941
+ label=f"DD start ({n_train})",
942
+ )
943
+
944
+ _set_eng_formatter(axi, "x", "s")
945
+ axi.set_xlabel("Time [s]")
946
+ axi.set_ylabel("Phase [rad]")
947
+ axi.set_title(title)
948
+ axi.legend(loc="best")
949
+
950
+ if show:
951
+ plt.show()
952
+ return None
953
+ return fig, axi