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,424 @@
1
+ """MAP Tikhonov carrier phase recovery with RTS/SSKF smoothers."""
2
+
3
+ import logging
4
+
5
+ import numpy as np
6
+
7
+ from ..backend import ArrayType, dispatch, to_device
8
+ from ..frequency import _modulation_power_m
9
+ from ..logger import logger
10
+ from .corrections import correct_cycle_slips
11
+
12
+ _NUMBA_RTS: dict = {}
13
+
14
+
15
+ def _get_numba_rts_smoother():
16
+ """JIT-compile and cache the Numba RTS-smoother kernel.
17
+
18
+ Returns
19
+ -------
20
+ callable
21
+ Numba-compiled ``_rts_loop``.
22
+ """
23
+ if "rts" not in _NUMBA_RTS:
24
+ import numba
25
+
26
+ @numba.njit(cache=True, fastmath=True, nogil=True)
27
+ def _rts_loop(phi_obs, sigma_p2, sigma_v2):
28
+ """Rauch-Tung-Striebel smoother - Numba inner kernel.
29
+
30
+ Parameters
31
+ ----------
32
+ phi_obs : (B,) float64
33
+ sigma_p2 : float64
34
+ sigma_v2 : float64
35
+
36
+ Returns
37
+ -------
38
+ (B,) float64
39
+ """
40
+ B = len(phi_obs)
41
+ x_filt = np.empty(B, dtype=np.float64)
42
+ P_filt = np.empty(B, dtype=np.float64)
43
+ x_pred = np.empty(B, dtype=np.float64)
44
+ P_pred = np.empty(B, dtype=np.float64)
45
+
46
+ x_filt[0] = phi_obs[0]
47
+ P_filt[0] = sigma_v2
48
+
49
+ for k in range(1, B):
50
+ x_pred[k] = x_filt[k - 1]
51
+ P_pred[k] = P_filt[k - 1] + sigma_p2
52
+ K = P_pred[k] / (P_pred[k] + sigma_v2)
53
+ x_filt[k] = x_pred[k] + K * (phi_obs[k] - x_pred[k])
54
+ P_filt[k] = (1.0 - K) * P_pred[k]
55
+
56
+ x_smooth = x_filt.copy()
57
+ for k in range(B - 2, -1, -1):
58
+ G = P_filt[k] / P_pred[k + 1]
59
+ x_smooth[k] = x_filt[k] + G * (x_smooth[k + 1] - x_pred[k + 1])
60
+
61
+ return x_smooth
62
+
63
+ _NUMBA_RTS["rts"] = _rts_loop
64
+
65
+ return _NUMBA_RTS["rts"]
66
+
67
+
68
+ def _rts_smoother_1d(
69
+ phi_obs: np.ndarray,
70
+ sigma_p2: float,
71
+ sigma_v2: float,
72
+ ) -> np.ndarray:
73
+ """Rauch-Tung-Striebel (RTS) Kalman smoother for a 1-D random-walk state.
74
+
75
+ Uses the Numba-compiled kernel (``_get_numba_rts_smoother``) when
76
+ available; falls back to a pure-Python loop otherwise. Always runs on
77
+ CPU - call with a NumPy array; the caller is responsible for
78
+ ``to_device`` conversion.
79
+
80
+ State model : x[k+1] = x[k] + w[k], w ~ N(0, sigma_p2)
81
+ Observation : y[k] = x[k] + v[k], v ~ N(0, sigma_v2)
82
+
83
+ Parameters
84
+ ----------
85
+ phi_obs : (B,) float64
86
+ Noisy block-phase observations in radians (e.g. from VV).
87
+ sigma_p2 : float
88
+ Process noise variance per block (Wiener phase noise increment).
89
+ sigma_v2 : float
90
+ Observation noise variance (VV estimator variance per block).
91
+
92
+ Returns
93
+ -------
94
+ (B,) float64
95
+ MAP-smoothed phase trajectory.
96
+ """
97
+ return _get_numba_rts_smoother()(phi_obs, float(sigma_p2), float(sigma_v2))
98
+
99
+
100
+ def _sskf_smoother_1d(
101
+ phi_obs: ArrayType,
102
+ sigma_p2: float,
103
+ sigma_v2: float,
104
+ sp,
105
+ xp,
106
+ ) -> ArrayType:
107
+ """Steady-state Kalman smoother via zero-phase IIR filter (filtfilt).
108
+
109
+ Approximates the RTS smoother by replacing the sequential Kalman
110
+ recurrence with a 1st-order IIR filter whose gain is solved analytically
111
+ from the discrete algebraic Riccati equation. The bidirectional
112
+ ``filtfilt`` call makes it equivalent to the RTS smoother in steady state.
113
+
114
+ Backend-aware: uses ``sp.signal.filtfilt`` where ``sp`` is
115
+ ``scipy`` (CPU) or ``cupyx.scipy`` (GPU) as returned by
116
+ ``dispatch``.
117
+
118
+ The approximation is excellent when ``B >> 1/K_∞``
119
+ (typically ``B > 20``). For ``B < 7`` (``filtfilt`` minimum), falls
120
+ back to the exact ``_rts_smoother_1d`` on CPU.
121
+
122
+ Parameters
123
+ ----------
124
+ phi_obs : (B,) float64, on the target device
125
+ Noisy block-phase observations in radians.
126
+ sigma_p2, sigma_v2 : float
127
+ Process and observation noise variances per block.
128
+ sp : module
129
+ ``scipy`` or ``cupyx.scipy``, from ``dispatch``.
130
+ xp : module
131
+ ``numpy`` or ``cupy``, from ``dispatch``.
132
+
133
+ Returns
134
+ -------
135
+ (B,) float64, same device as ``phi_obs``.
136
+ """
137
+ # filtfilt requires at least padlen * 2 + 1 samples; padlen = 3 * max(len(b), len(a)) = 6
138
+ if len(phi_obs) < 7:
139
+ phi_np = to_device(phi_obs, "cpu")
140
+ return xp.asarray(_rts_smoother_1d(phi_np, sigma_p2, sigma_v2))
141
+
142
+ # Steady-state prediction error covariance from discrete Riccati equation:
143
+ # p² - σ_p²·p - σ_p²·σ_v² = 0 -> p = (σ_p² + √(σ_p⁴ + 4σ_p²σ_v²)) / 2
144
+ p_ss = (sigma_p2 + float(np.sqrt(sigma_p2**2 + 4.0 * sigma_p2 * sigma_v2))) / 2.0
145
+ K_ss = p_ss / (p_ss + sigma_v2)
146
+
147
+ # Forward IIR: y[k] = (1-K)·y[k-1] + K·x[k]
148
+ # H(z) = K / (1 - (1-K)·z⁻¹)
149
+ # filtfilt applies forward + backward -> zero-phase, ≡ RTS smoother at
150
+ # steady state.
151
+ b = [K_ss]
152
+ a = [1.0, -(1.0 - K_ss)]
153
+ return sp.signal.filtfilt(b, a, phi_obs)
154
+
155
+
156
+ def recover_carrier_phase_tikhonov(
157
+ symbols: ArrayType,
158
+ modulation: str,
159
+ order: int,
160
+ linewidth_symbol_periods: float,
161
+ block_size: int = 32,
162
+ snr_db: float | None = None,
163
+ method: str = "exact",
164
+ joint_channels: bool = False,
165
+ cycle_slip_correction: bool = False,
166
+ cycle_slip_history: int = 100,
167
+ cycle_slip_threshold: float = np.pi / 4,
168
+ debug_plot: bool = False,
169
+ ) -> ArrayType:
170
+ r"""
171
+ Carrier phase recovery via MAP estimation with a Tikhonov/Wiener phase
172
+ noise prior.
173
+
174
+ Extends the Viterbi-Viterbi block estimator with a Kalman smoother
175
+ matched to the laser phase noise statistics. Two smoother backends are
176
+ available via ``method``:
177
+
178
+ * ``'exact'`` - full Rauch-Tung-Striebel (RTS) smoother; Numba-compiled.
179
+ Exact for all sequence lengths; runs on CPU.
180
+ * ``'sskf'`` - steady-state Kalman filter approximation via zero-phase
181
+ IIR (``filtfilt``); backend-aware (stays on GPU when input is on GPU).
182
+ Approximation holds for ``N_blocks >> 1/K_∞`` (~20+ blocks typical).
183
+
184
+ Parameters
185
+ ----------
186
+ symbols : array_like
187
+ 1-SPS complex symbols after matched filter and FOE.
188
+ Shape: ``(N,)`` or ``(C, N)``.
189
+ modulation : str
190
+ Modulation scheme (case-insensitive): ``'psk'``, ``'qam'``, etc.
191
+ order : int
192
+ Modulation order.
193
+ linewidth_symbol_periods : float
194
+ Combined linewidth-symbol-time product delta_nu * T_s.
195
+ Typical values: ``1e-5`` (narrow laser, 32 GBd), ``5e-4`` (wide
196
+ laser / high baud rate). Sets the Kalman process noise variance:
197
+ sigma_p^2 = 2*pi * delta_nu * T_s * N_b.
198
+ block_size : int, default 32
199
+ Symbols per VV estimation block. Same trade-off as for
200
+ ``recover_carrier_phase_viterbi_viterbi``.
201
+ snr_db : float or None, default None
202
+ Per-symbol SNR in dB. Used to compute the VV observation noise
203
+ variance sigma_v^2 ≈ 1 / (M^2 * SNR * N_b).
204
+ If ``None``, defaults to 20 dB with a warning - provide the actual
205
+ operating SNR for the optimal smoother bandwidth.
206
+ method : {'exact', 'sskf'}, default 'exact'
207
+ Smoother implementation:
208
+
209
+ * ``'exact'``: full RTS smoother (``_rts_smoother_1d``); Numba
210
+ kernel when available. Sequential CPU recurrence; exact for any
211
+ ``N_blocks``. On GPU inputs this forces a full device-to-host
212
+ transfer of the block-phase trajectory (and back), stalling the
213
+ GPU pipeline - prefer ``'sskf'`` for GPU-resident signals.
214
+ * ``'sskf'``: steady-state approximation via ``filtfilt``
215
+ (``_sskf_smoother_1d``); runs on the input device (GPU-native
216
+ when data is on GPU, no host transfer). Excellent for
217
+ ``N_blocks ≥ 20``; for ``N_blocks < 7`` silently falls back to
218
+ ``'exact'``.
219
+ joint_channels : bool, default False
220
+ For MIMO inputs (C > 1): if ``True``, sum the M-th-power block
221
+ phasors across all channels before the VV phase extraction and
222
+ Kalman smoother. The single smoothed trajectory is broadcast to
223
+ all C output rows. Reduces variance by ~√C for shared-LO systems.
224
+ cycle_slip_correction : bool, default False
225
+ If ``True``, apply cycle-slip detection and correction
226
+ (``correct_cycle_slips``) after the Kalman smoother, before
227
+ interpolation.
228
+ cycle_slip_history : int, default 100
229
+ ``history_length`` passed to ``correct_cycle_slips``.
230
+ cycle_slip_threshold : float, default π/4
231
+ ``threshold`` passed to ``correct_cycle_slips`` (radians).
232
+ debug_plot : bool, default False
233
+ If ``True``, opens a diagnostic figure showing the per-symbol phase
234
+ trajectory with the Kalman-smoothed block phases.
235
+
236
+ Returns
237
+ -------
238
+ array_like
239
+ Per-symbol phase estimate in radians. Shape matches ``symbols``.
240
+ Same backend as input.
241
+
242
+ Notes
243
+ -----
244
+ VV block phases are Kalman-smoothed with sigma_p^2 = 2*pi*linewidth*T_s*N_b
245
+ and sigma_v^2 ≈ 1/(M^2 * SNR * N_b), then interpolated to per-symbol
246
+ resolution. A residual 2*pi/M ambiguity always remains.
247
+ """
248
+ if method not in ("exact", "sskf"):
249
+ raise ValueError(f"Unknown method {method!r}. Choose 'exact' or 'sskf'.")
250
+
251
+ symbols, xp, sp = dispatch(symbols)
252
+ was_1d = symbols.ndim == 1
253
+ if was_1d:
254
+ symbols = symbols[None, :]
255
+ C, N = symbols.shape
256
+
257
+ M = _modulation_power_m(modulation, order)
258
+
259
+ N_trunc = (N // block_size) * block_size
260
+ N_blocks = N_trunc // block_size
261
+
262
+ if N_blocks == 0:
263
+ raise ValueError(
264
+ f"Signal length {N} is shorter than block_size={block_size}. "
265
+ "Reduce block_size or use a longer symbol sequence."
266
+ )
267
+
268
+ # Same data-residual constraint as VV: for QAM with order > 4 the M-th power
269
+ # does not cancel per symbol. Block phase variance can exceed π/M before the
270
+ # Kalman smoother is applied, causing unwrap slips that the smoother cannot fix.
271
+ if "qam" in modulation.lower() and order > 4:
272
+ _min_bs = max(8, 4 * int(np.ceil(order**0.5)))
273
+ if block_size < _min_bs:
274
+ logger.warning(
275
+ "CPR (Tikhonov): block_size=%s is too small for %s-QAM. Block phases are estimated via Viterbi-Viterbi; the data-residual constraint is identical - see recover_carrier_phase_viterbi_viterbi. Recommended minimum for %s-QAM: block_size ≥ %s.",
276
+ block_size,
277
+ order,
278
+ order,
279
+ _min_bs,
280
+ )
281
+
282
+ # Smoother noise parameters
283
+ if snr_db is None:
284
+ logger.warning(
285
+ "CPR (Tikhonov): snr_db not provided - defaulting to 20 dB. "
286
+ "Pass the operating SNR for the optimal smoother bandwidth."
287
+ )
288
+ snr_lin = 100.0 # 20 dB default
289
+ else:
290
+ snr_lin = 10.0 ** (snr_db / 10.0)
291
+
292
+ sigma_p2 = float(2.0 * np.pi * linewidth_symbol_periods * block_size)
293
+ sigma_v2 = float(1.0 / (M**2 * snr_lin * block_size))
294
+
295
+ # VV block phase estimation with unit-circle normalisation for QAM
296
+ blocks = symbols[:, :N_trunc].reshape(C, N_blocks, block_size)
297
+ blocks_c = blocks.astype(
298
+ xp.complex128 if blocks.dtype == xp.complex64 else blocks.dtype
299
+ )
300
+ if "qam" in modulation.lower():
301
+ mag = xp.abs(blocks_c)
302
+ blocks_c = blocks_c / xp.maximum(mag, 1e-15 * xp.max(mag))
303
+
304
+ S_b = xp.sum(blocks_c**M, axis=-1) # (C, N_blocks)
305
+
306
+ block_centers = xp.arange(N_blocks, dtype=xp.float64) * block_size + block_size / 2
307
+ all_positions = xp.arange(N, dtype=xp.float64)
308
+ phi_full = xp.zeros((C, N), dtype=xp.float64)
309
+
310
+ if joint_channels and C > 1:
311
+ # Sum M-th-power phasors -> single VV estimate -> single Kalman pass
312
+ S_b_joint = xp.sum(S_b, axis=0) # (N_blocks,)
313
+ phi_raw_joint = xp.angle(S_b_joint) / M
314
+ phi_u_joint = xp.unwrap((phi_raw_joint * M).astype(xp.float64)) / M
315
+ if "qam" in modulation.lower():
316
+ phi_u_joint = phi_u_joint - (np.pi / M)
317
+
318
+ # Kalman smoother on the joint trajectory
319
+ if method == "exact":
320
+ phi_u_joint_np = to_device(phi_u_joint, "cpu")
321
+ phi_smooth_joint_np = _rts_smoother_1d(phi_u_joint_np, sigma_p2, sigma_v2)
322
+ phi_smooth_joint = xp.asarray(phi_smooth_joint_np)
323
+ else:
324
+ phi_smooth_joint = _sskf_smoother_1d(
325
+ phi_u_joint, sigma_p2, sigma_v2, sp, xp
326
+ )
327
+ phi_smooth_joint_np = to_device(phi_smooth_joint, "cpu")
328
+
329
+ if cycle_slip_correction:
330
+ phi_smooth_joint_np = correct_cycle_slips(
331
+ to_device(phi_smooth_joint, "cpu"),
332
+ 4,
333
+ cycle_slip_history,
334
+ cycle_slip_threshold,
335
+ )
336
+ phi_smooth_joint = xp.asarray(phi_smooth_joint_np)
337
+
338
+ phi_interp = xp.interp(all_positions, block_centers, phi_smooth_joint)
339
+ for ch in range(C):
340
+ phi_full[ch] = phi_interp
341
+ phi_smooth_np = np.tile(to_device(phi_smooth_joint, "cpu"), (C, 1))
342
+ else:
343
+ phi_raw = xp.angle(S_b) / M
344
+ phi_u = (
345
+ xp.unwrap((phi_raw * M).astype(xp.float64), axis=-1) / M
346
+ ) # (C, N_blocks)
347
+
348
+ if "qam" in modulation.lower():
349
+ phi_u = phi_u - (np.pi / M)
350
+
351
+ if C > 1:
352
+ # All per-channel means on device, one batched D2H, vectorized shift
353
+ # (instead of one float() sync + one rounding per channel).
354
+ diffs_np = to_device(xp.mean(phi_u[1:] - phi_u[0:1], axis=-1), "cpu")
355
+ k_np = np.round(diffs_np * M / (2 * np.pi))
356
+ phi_u[1:] = phi_u[1:] - xp.asarray(k_np)[:, None] * (2 * np.pi / M)
357
+
358
+ # Kalman smoother - dispatch on method
359
+ if method == "exact":
360
+ phi_u_np = to_device(phi_u, "cpu") # (C, N_blocks) float64
361
+ phi_smooth_np = np.empty_like(phi_u_np)
362
+ for ch in range(C):
363
+ phi_smooth_np[ch] = _rts_smoother_1d(phi_u_np[ch], sigma_p2, sigma_v2)
364
+ phi_smooth = xp.asarray(phi_smooth_np)
365
+ else: # method == "sskf"
366
+ phi_smooth = xp.empty_like(phi_u)
367
+ for ch in range(C):
368
+ phi_smooth[ch] = _sskf_smoother_1d(
369
+ phi_u[ch], sigma_p2, sigma_v2, sp, xp
370
+ )
371
+ phi_smooth_np = to_device(phi_smooth, "cpu")
372
+
373
+ for ch in range(C):
374
+ phi_s_ch = phi_smooth[ch]
375
+ if cycle_slip_correction:
376
+ phi_s_ch_np = correct_cycle_slips(
377
+ to_device(phi_s_ch, "cpu"),
378
+ 4,
379
+ cycle_slip_history,
380
+ cycle_slip_threshold,
381
+ )
382
+ phi_s_ch = xp.asarray(phi_s_ch_np)
383
+ phi_smooth_np[ch] = to_device(phi_s_ch, "cpu")
384
+ phi_full[ch] = xp.interp(all_positions, block_centers, phi_s_ch)
385
+
386
+ # Host copy of the trajectory is needed only for the INFO summary and the
387
+ # optional debug plot; skip the transfer + reductions otherwise (the device
388
+ # phi_full drives the actual correction and is what gets returned).
389
+ _want_log = logger.isEnabledFor(logging.INFO)
390
+ if _want_log or debug_plot:
391
+ phi_full_np = to_device(phi_full, "cpu")
392
+ if _want_log:
393
+ phi_mean_deg = float(np.mean(phi_full_np)) * 180.0 / np.pi
394
+ phi_std_deg = float(np.std(phi_full_np)) * 180.0 / np.pi
395
+ mode_str = "joint" if (joint_channels and C > 1) else "independent"
396
+ logger.info(
397
+ "CPR (Tikhonov-%s, M=%s, %s): phase mean=%.2f°, std=%.2f° [%s blocks x %s, σ_p²=%.2e, σ_v²=%.2e, C=%s, cycle_slip_correction=%s]",
398
+ method.upper(),
399
+ M,
400
+ mode_str,
401
+ phi_mean_deg,
402
+ phi_std_deg,
403
+ N_blocks,
404
+ block_size,
405
+ sigma_p2,
406
+ sigma_v2,
407
+ C,
408
+ cycle_slip_correction,
409
+ )
410
+
411
+ if debug_plot:
412
+ from .. import plotting as _plotting
413
+
414
+ _plotting.plot_carrier_phase_trajectory(
415
+ phi_full=phi_full_np,
416
+ block_centers=to_device(block_centers, "cpu"),
417
+ phi_blocks=phi_smooth_np,
418
+ show=True,
419
+ title=f"CPR - Tikhonov-{method.upper()}",
420
+ )
421
+
422
+ if was_1d:
423
+ return phi_full[0]
424
+ return phi_full
@@ -0,0 +1,227 @@
1
+ """Viterbi-Viterbi (V&V) carrier phase recovery."""
2
+
3
+ import logging
4
+
5
+ import numpy as np
6
+
7
+ from ..backend import ArrayType, dispatch, to_device
8
+ from ..frequency import _modulation_power_m
9
+ from ..logger import logger
10
+ from .corrections import correct_cycle_slips
11
+
12
+
13
+ def recover_carrier_phase_viterbi_viterbi(
14
+ symbols: ArrayType,
15
+ modulation: str,
16
+ order: int,
17
+ block_size: int = 32,
18
+ joint_channels: bool = False,
19
+ cycle_slip_correction: bool = False,
20
+ cycle_slip_history: int = 100,
21
+ cycle_slip_threshold: float = np.pi / 4,
22
+ debug_plot: bool = False,
23
+ ) -> ArrayType:
24
+ """
25
+ Carrier phase recovery via the Viterbi-Viterbi (M-th power) algorithm.
26
+
27
+ Block-based blind phase estimation for PSK and QAM symbols. Raises each
28
+ block of symbols to the M-th power to remove modulation, extracts the
29
+ block phase, resolves the M-fold ambiguity by unwrapping, then
30
+ interpolates to per-symbol resolution.
31
+
32
+ Parameters
33
+ ----------
34
+ symbols : array_like
35
+ 1-SPS complex symbols after matched filter. Shape: (N,) or (C, N).
36
+ modulation : str
37
+ Modulation scheme (case-insensitive): 'psk', 'qam', etc.
38
+ order : int
39
+ Modulation order.
40
+ block_size : int, default 32
41
+ Number of symbols per estimation block. Larger blocks reduce
42
+ variance but reduce tracking bandwidth for fast phase noise.
43
+ Typical range: 16-128 for QAM; as low as 1 for PSK (data cancels
44
+ exactly in the M-th power for M-PSK constellations).
45
+ joint_channels : bool, default False
46
+ For MIMO inputs (C > 1): if ``True``, sum the M-th-power block
47
+ phasors ``S_b`` across all channels before phase extraction.
48
+ The resulting single trajectory is broadcast to all C output rows.
49
+ Reduces variance by ~√C for shared-LO systems. SISO-safe.
50
+ cycle_slip_correction : bool, default False
51
+ If ``True``, apply cycle-slip detection and correction
52
+ (``correct_cycle_slips``) after M-fold unwrap, before
53
+ interpolation.
54
+ cycle_slip_history : int, default 100
55
+ ``history_length`` passed to ``correct_cycle_slips``.
56
+ cycle_slip_threshold : float, default π/4
57
+ ``threshold`` passed to ``correct_cycle_slips`` (radians).
58
+ debug_plot : bool, default False
59
+ If ``True``, opens a diagnostic figure showing the per-symbol phase
60
+ trajectory alongside the block-phase estimates.
61
+
62
+ Returns
63
+ -------
64
+ array_like
65
+ Per-symbol phase estimate in radians. Shape matches ``symbols``.
66
+ Same backend as input.
67
+
68
+ Notes
69
+ -----
70
+ Each block: S_b = sum s[n]^M, phi_hat_b = angle(S_b) / M. Block phases
71
+ are M-fold unwrapped; a global 2*pi/M ambiguity always remains.
72
+
73
+ For QAM with order > 4, block averaging suppresses M-th-power data residuals;
74
+ minimum reliable block_size scales as ~4*ceil(sqrt(order)). For high phase
75
+ noise prefer ``recover_carrier_phase_bps`` (no unwrap required).
76
+ """
77
+ symbols, xp, _ = dispatch(symbols)
78
+ was_1d = symbols.ndim == 1
79
+ if was_1d:
80
+ symbols = symbols[None, :] # (1, N)
81
+ C, N = symbols.shape
82
+
83
+ M = _modulation_power_m(modulation, order)
84
+
85
+ N_trunc = (N // block_size) * block_size
86
+ N_blocks = N_trunc // block_size
87
+
88
+ if N_blocks == 0:
89
+ raise ValueError(
90
+ f"Signal length {N} is shorter than block_size={block_size}. "
91
+ "Reduce block_size or use a longer symbol sequence."
92
+ )
93
+
94
+ # For QAM with order > 4 the M-th power of individual symbols does NOT cancel
95
+ # the data modulation (unlike PSK, where every M-PSK point gives (c/|c|)^M = 1).
96
+ # Sufficient block averaging is required so that the block-phase variance stays
97
+ # below the π/M unwrap threshold. The practical minimum scales as 4·ceil(√order).
98
+ if "qam" in modulation.lower() and order > 4:
99
+ _min_bs = max(8, 4 * int(np.ceil(order**0.5)))
100
+ if block_size < _min_bs:
101
+ logger.warning(
102
+ "CPR (VV): block_size=%s is too small for %s-QAM. Individual QAM symbols' M-th powers do not cancel the data modulation; insufficient averaging causes block-phase variance that exceeds the π/M unwrap threshold, producing persistent 2π/M phase slips. Recommended minimum for %s-QAM: block_size ≥ %s.",
103
+ block_size,
104
+ order,
105
+ order,
106
+ _min_bs,
107
+ )
108
+
109
+ # Reshape for block processing: (C, N_blocks, block_size).
110
+ # Promote to complex128 for the M-th power - identical to estimate_frequency_offset_mth_power.
111
+ # On GPU, complex64^4 loses precision near the ±π/M unwrap boundary, causing
112
+ # spurious branch flips for high-order QAM with small block sizes.
113
+ blocks = symbols[:, :N_trunc].reshape(C, N_blocks, block_size)
114
+ blocks_c = blocks.astype(
115
+ xp.complex128 if blocks.dtype == xp.complex64 else blocks.dtype
116
+ )
117
+
118
+ # For QAM, project to unit circle before the M-th power (normalized VV).
119
+ # This removes outer-ring amplitude dominance and makes the π/M QAM bias
120
+ # correction exact (by the 4-fold rotational symmetry of the constellation).
121
+ # PSK is already constant-modulus; normalization is a no-op.
122
+ if "qam" in modulation.lower():
123
+ mag = xp.abs(blocks_c)
124
+ blocks_c = blocks_c / xp.maximum(mag, 1e-15 * xp.max(mag))
125
+
126
+ S_b = xp.sum(blocks_c**M, axis=-1) # (C, N_blocks)
127
+
128
+ # Block centre positions for interpolation (uniform spacing = block_size)
129
+ block_centers = xp.arange(N_blocks, dtype=xp.float64) * block_size + block_size / 2
130
+ all_positions = xp.arange(N, dtype=xp.float64)
131
+
132
+ phi_full = xp.zeros((C, N), dtype=xp.float64)
133
+ phi_blocks_out = xp.zeros((C, N_blocks), dtype=xp.float64)
134
+
135
+ if joint_channels and C > 1:
136
+ # Sum M-th-power phasors across channels -> single block-phase trajectory
137
+ S_b_joint = xp.sum(S_b, axis=0) # (N_blocks,)
138
+ phi_raw_joint = xp.angle(S_b_joint) / M
139
+ phi_u_joint = xp.unwrap((phi_raw_joint * M).astype(xp.float64)) / M
140
+ if "qam" in modulation.lower():
141
+ phi_u_joint = phi_u_joint - (np.pi / M)
142
+ if cycle_slip_correction:
143
+ phi_u_joint_np = correct_cycle_slips(
144
+ to_device(phi_u_joint, "cpu"),
145
+ 4,
146
+ cycle_slip_history,
147
+ cycle_slip_threshold,
148
+ )
149
+ phi_u_joint = xp.asarray(phi_u_joint_np)
150
+ phi_interp = xp.interp(all_positions, block_centers, phi_u_joint)
151
+ for ch in range(C):
152
+ phi_full[ch] = phi_interp
153
+ phi_blocks_out[ch] = phi_u_joint
154
+ else:
155
+ # Raw block phase in [-π/M, π/M)
156
+ phi_raw = xp.angle(S_b) / M # (C, N_blocks)
157
+
158
+ # M-fold unwrap: scale into 2π domain, unwrap, re-scale back.
159
+ # Cast to float64 before unwrap - cp.unwrap preserves input dtype so float32
160
+ # would lose precision during the discontinuity test (diff vs 2π threshold).
161
+ phi_u = (
162
+ xp.unwrap((phi_raw * M).astype(xp.float64), axis=-1) / M
163
+ ) # (C, N_blocks)
164
+
165
+ # QAM bias correction.
166
+ if "qam" in modulation.lower():
167
+ phi_u = phi_u - (np.pi / M)
168
+
169
+ # MIMO M-fold alignment: align every channel to channel 0's branch.
170
+ # Skipped in joint mode (all channels share the same trajectory).
171
+ if C > 1:
172
+ # All per-channel means on device, one batched D2H, vectorized shift
173
+ # (instead of one float() sync + one rounding per channel).
174
+ diffs_np = to_device(xp.mean(phi_u[1:] - phi_u[0:1], axis=-1), "cpu")
175
+ k_np = np.round(diffs_np * M / (2 * np.pi))
176
+ phi_u[1:] = phi_u[1:] - xp.asarray(k_np)[:, None] * (2 * np.pi / M)
177
+
178
+ # xp.interp is 1D-only; loop over C channels.
179
+ for ch in range(C):
180
+ phi_u_ch = phi_u[ch]
181
+ if cycle_slip_correction:
182
+ phi_u_ch_np = correct_cycle_slips(
183
+ to_device(phi_u_ch, "cpu"),
184
+ 4,
185
+ cycle_slip_history,
186
+ cycle_slip_threshold,
187
+ )
188
+ phi_u_ch = xp.asarray(phi_u_ch_np)
189
+ phi_full[ch] = xp.interp(all_positions, block_centers, phi_u_ch)
190
+ phi_blocks_out[ch] = phi_u_ch
191
+
192
+ # Host copy of the trajectory is needed only for the INFO summary and the
193
+ # optional debug plot; skip the transfer + reductions otherwise (the device
194
+ # phi_full drives the actual correction and is what gets returned).
195
+ _want_log = logger.isEnabledFor(logging.INFO)
196
+ if _want_log or debug_plot:
197
+ phi_full_np = to_device(phi_full, "cpu")
198
+ if _want_log:
199
+ phi_mean_deg = float(np.mean(phi_full_np)) * 180.0 / np.pi
200
+ phi_std_deg = float(np.std(phi_full_np)) * 180.0 / np.pi
201
+ mode_str = "joint" if (joint_channels and C > 1) else "independent"
202
+ logger.info(
203
+ "CPR (Viterbi-Viterbi, M=%s, %s): phase mean=%.2f°, std=%.2f° [%s blocks x %s symbols, C=%s, cycle_slip_correction=%s]",
204
+ M,
205
+ mode_str,
206
+ phi_mean_deg,
207
+ phi_std_deg,
208
+ N_blocks,
209
+ block_size,
210
+ C,
211
+ cycle_slip_correction,
212
+ )
213
+
214
+ if debug_plot:
215
+ from .. import plotting as _plotting
216
+
217
+ _plotting.plot_carrier_phase_trajectory(
218
+ phi_full=phi_full_np,
219
+ block_centers=to_device(block_centers, "cpu"),
220
+ phi_blocks=to_device(phi_blocks_out, "cpu"),
221
+ show=True,
222
+ title="CPR - Viterbi-Viterbi",
223
+ )
224
+
225
+ if was_1d:
226
+ return phi_full[0]
227
+ return phi_full